Personalize this lesson
Adapt explanations and teaching visuals to your background and preferred voice.
Sequence models compress an ordered history into a compact hidden state. An inspection vision pipeline makes the same move on image data: taking a high-dimensional screenshot crop and squeezing it down to a handful of numbers.
Can two scalar numbers preserve a bright 2 by 2 alert badge on a 5 by 5 grayscale screenshot? And if you hand the decoder two numbers drawn at random instead of an encoded crop, should it know what to draw? Compression and generation impose different constraints. An autoencoder learns to reconstruct its input from a latent code, which we make compact here.[1] A variational autoencoder (VAE) learns a probabilistic latent model with a prior we can sample for generation. Its objective connects reconstruction and prior sampling, but useful samples still depend on the model, data, and training.[2]
Compress one screenshot crop
Start with a tiny crop. Each grayscale number is between 0 (dark) and 1 (bright). A real screenshot has many more pixels, but a 5 by 5 crop lets us inspect every single number. The bright 2 by 2 block is the highlighted region. Before looking at encoder maps, predict: which pixels should one latent coordinate watch, and which should it ignore?
![An end-to-end autoencoder hourglass: a 5 by 5 grayscale crop (25 pixels) is compressed by two encoder weight maps into a two-number latent bottleneck z = [0.888, 0.170], then expanded by a decoder into a reconstructed crop x_hat with reconstruction MSE of 0.0027.](/cdn/content-image/preparation/autoencoders-vaes-generative-modeling-basics/illustrations/_generated/grayscale_patch_dark.png?v=5614ed37d27d)
An autoencoder has two learned functions:
The encoder maps the input to the latent code . Meanwhile, the decoder maps that code to a reconstruction . Here and denote their learned weights. A narrow creates the hourglass bottleneck: there's no direct wire carrying all 25 pixels to the decoder. This forces the network to discover compact summaries, but doesn't prevent a large network from memorizing a tiny dataset.
For this crop, the first encoder map averages the four bright center pixels and the second averages two lower pixels:
Those aren't learned weights. They're a hand-designed encoder, so the bottleneck arithmetic is visible before we ask optimization to discover one.
For normalized pixels, mean-squared error (MSE) is a simple reconstruction measure:
is the number of pixels, here 25. Each term is one pixel's squared error. The next snippet runs that same encoder, then a matching crafted decoder, so you can see the 25-to-2 shrink and the reconstruction error together.
1import numpy as np
2
3x = np.array([
4 0.08, 0.12, 0.09, 0.11, 0.07,
5 0.15, 0.85, 0.88, 0.14, 0.10,
6 0.11, 0.90, 0.92, 0.13, 0.09,
7 0.07, 0.18, 0.16, 0.12, 0.08,
8 0.06, 0.09, 0.10, 0.07, 0.05,
9], dtype=np.float32)
10
11center = [6, 7, 11, 12]
12lower = [16, 17]
13W_encoder = np.zeros((25, 2), dtype=np.float32)
14W_encoder[center, 0] = 0.25
15W_encoder[lower, 1] = 0.50
16
17W_decoder = np.zeros((2, 25), dtype=np.float32)
18W_decoder[0, center] = 1.0
19W_decoder[0, [5, 8, 10, 13]] = 0.12
20W_decoder[1, lower] = 1.0
21W_decoder[1, [15, 18, 21, 22]] = 0.35
22
23z = x @ W_encoder
24x_hat = np.clip(z @ W_decoder + 0.08, 0.0, 1.0)
25mse = np.mean((x - x_hat) ** 2)
26
27print("input numbers:", x.size)
28print("latent numbers:", z.size, np.round(z, 3))
29print("reconstruction MSE:", round(float(mse), 4))1input numbers: 25
2latent numbers: 2 [0.888 0.17 ]
3reconstruction MSE: 0.0027The input shrinks from 25 numbers to 2. This crafted decoder does well on this one pattern because we built its weights around it. Notice what it loses: rearranging the four selected center pixels while preserving their average leaves unchanged. The decoder then produces the same reconstruction for different inputs. To handle a family of crops, learn the weights instead of selecting them by hand.
Learn the bottleneck
The next program makes noisy variations of the same crop and trains an autoencoder in PyTorch. The training loop uses gradients to update weights with an optimizer, as in the backpropagation lesson. Compare its MSE with a deliberately uninformative baseline that returns the average crop for every input.
1import torch
2from torch import nn
3import torch.nn.functional as F
4
5torch.manual_seed(3)
6base = torch.tensor([
7 0.08, 0.12, 0.09, 0.11, 0.07,
8 0.15, 0.85, 0.88, 0.14, 0.10,
9 0.11, 0.90, 0.92, 0.13, 0.09,
10 0.07, 0.18, 0.16, 0.12, 0.08,
11 0.06, 0.09, 0.10, 0.07, 0.05,
12])
13patches = torch.clamp(base + 0.025 * torch.randn(48, 25), 0.0, 1.0)
14
15encoder = nn.Sequential(nn.Linear(25, 8), nn.Tanh(), nn.Linear(8, 2))
16decoder = nn.Sequential(nn.Linear(2, 8), nn.Tanh(), nn.Linear(8, 25), nn.Sigmoid())
17optimizer = torch.optim.Adam(
18 list(encoder.parameters()) + list(decoder.parameters()), lr=0.03
19)
20
21with torch.no_grad():
22 before = F.mse_loss(decoder(encoder(patches)), patches).item()
23
24for _ in range(250):
25 reconstruction = decoder(encoder(patches))
26 loss = F.mse_loss(reconstruction, patches)
27 optimizer.zero_grad()
28 loss.backward()
29 optimizer.step()
30
31with torch.no_grad():
32 after = F.mse_loss(decoder(encoder(patches)), patches).item()
33 latent = encoder(patches[:1])
34 mean_crop = patches.mean(dim=0, keepdim=True).expand_as(patches)
35 baseline = F.mse_loss(mean_crop, patches).item()
36
37print(f"MSE before training: {before:.4f}")
38print(f"MSE after training: {after:.4f}")
39print(f"constant-crop MSE: {baseline:.4f}")
40print("latent shape:", tuple(latent.shape))1MSE before training: 0.1537
2MSE after training: 0.0006
3constant-crop MSE: 0.0006
4latent shape: (1, 2)The trained MSE is close to the constant-crop baseline. These nearly identical inputs don't require an informative code: a decoder can do well by remembering their shared pattern. Optimization works as expected, but this test alone doesn't prove the latent captures useful variation. A solid evaluation needs distinct patterns, held-out crops, and a verification that perturbing the latent changes the reconstruction.
Even an autoencoder that reconstructs many distinct patterns accurately can't decode arbitrary random pairs of numbers. Its reconstruction loss only penalizes codes produced by the encoder for actual training inputs:
The encoder is free to arrange those codes without matching a chosen prior. It might place center-highlight crops near and lower-strip crops near . A standard-normal draw can then fall in a region poorly represented by encoded training crops, where reconstruction loss doesn't directly assess the decoder. Its output there may be useful, blurry, or implausible; this objective doesn't tell us which. Neither separated clusters nor failed random samples are inevitable. A separate model could also learn a distribution over deterministic codes.
| Model question | Deterministic autoencoder answer |
|---|---|
| What does the encoder output? | One point for each input |
| What is optimized? | Reconstruction of encoded training examples |
| Can you decode a known encoded crop? | Yes, if reconstruction training worked |
| Is trained as an input source? | No; its relation to encoded crops is unconstrained by plain reconstruction loss |

Why isn't a low reconstruction loss enough to make a deterministic autoencoder a generator?
Answer
The reconstruction loss evaluates codes produced by the encoder for training-like inputs. It doesn't require the decoder to behave sensibly on points sampled from a chosen prior such as a standard normal distribution.
A VAE learns distributions over codes
A VAE starts with a generative story: draw a latent from a prior , then draw an image from a decoder distribution . For an observed crop, inference runs that story backward: which latent values could have produced it? Instead of returning one code, the encoder predicts an approximate posterior, written , over those possible values.
The word approximate matters: computing the exact posterior is generally impractical, so the encoder learns a tractable stand-in. This approximation method is variational inference. Here it's amortized: one encoder predicts posterior parameters for every crop instead of solving a fresh optimization problem for each input.
For the common diagonal-Gaussian case, it predicts a mean vector and a log-variance vector :
The model also declares a simple prior, usually:
The posterior describes codes for one observed crop. The prior is the distribution we sample from when no input crop is provided. In , each coordinate has mean zero and variance one, and the coordinates are independent. A diagonal posterior similarly has independent coordinates conditional on this crop, though its means and variances depend on the input. Mixing the posteriors of different crops needn't preserve that independence in the aggregate codes.
Why predict logvar rather than a raw standard deviation? An unconstrained linear layer can output negative values, while a standard deviation must be positive. Mathematically, exponentiating half a finite log variance makes a positive scale. Floating-point exponentiation can still underflow to zero or overflow to infinity. A positive transform such as softplus is another parameterization; log variance is convenient, not mandatory.
1import numpy as np
2
3logvar = np.array([-0.70, 0.10])
4sigma = np.exp(0.5 * logvar)
5variance = np.exp(logvar)
6
7print("log variance:", logvar)
8print("variance:", np.round(variance, 3))
9print("standard deviation:", np.round(sigma, 3))
10print("all sigma positive:", bool(np.all(sigma > 0)))1log variance: [-0.7 0.1]
2variance: [0.497 1.105]
3standard deviation: [0.705 1.051]
4all sigma positive: TrueThose two vectors define a Gaussian for one crop. Next we keep that Gaussian near a sampleable prior, then we'll sample from it without blocking gradients.
Keep posteriors near a sampleable prior
If the encoder may put each crop's distribution anywhere it likes, prior samples still won't reliably land where the decoder has trained. Before writing the penalty, predict what it should do: it should be zero when the posterior matches the prior, and grow when means or variances move away.
VAEs measure that mismatch with Kullback-Leibler (KL) divergence. KL isn't a symmetric distance: direction matters. For a diagonal Gaussian posterior and standard-normal prior, the closed form is:[2]
The penalty is zero when and . Moving means away from zero or making variances very different from one increases it.
Use the same two-dimensional posterior we'll sample from next: and . One coordinate at a time:
The next snippet checks that closed form on three posteriors, including this one.
1import numpy as np
2
3def kl_to_standard_normal(mu, logvar):
4 # expm1(logvar) computes exp(logvar) - 1 accurately near zero.
5 return float(0.5 * np.sum(mu**2 + np.expm1(logvar) - logvar))
6
7posteriors = {
8 "matches prior": (np.array([0.0, 0.0]), np.array([0.0, 0.0])),
9 "small offset": (np.array([0.4, -0.2]), np.array([-0.7, 0.1])),
10 "far offset": (np.array([2.0, -1.5]), np.array([-1.5, 1.2])),
11}
12
13for name, (mu, logvar) in posteriors.items():
14 print(f"{name:13} KL = {kl_to_standard_normal(mu, logvar):.4f}")1matches prior KL = 0.0000
2small offset KL = 0.2009
3far offset KL = 4.0466KL pressure isn't free. If it dominates reconstruction, all inputs can be encoded nearly like the prior and the decoder may ignore . We'll catch that failure after the training loop. First we need a way to sample from without blocking gradients to and .
Sample while gradients still flow
Kingma and Welling reparameterize a Gaussian draw so reconstruction gradients can reach encoder parameters.[2] Writing specifies a distribution, not how a sampling API connects to autograd. PyTorch's ordinary Normal.sample() supplies no pathwise gradient to its mean or scale. Reparameterization makes that path explicit. Other estimators, including score-function gradients, can train stochastic models without differentiating the sampled value directly.[3]
The reparameterization trick moves all stochasticity into an independent noise variable drawn from a standard normal distribution, then applies an algebraic transformation:
During the backward pass, is treated as a fixed constant tensor. The symbol means elementwise multiplication. For each coordinate, the partial derivatives are straightforward:
Reconstruction gradients can then reach , , and the encoder weights . PyTorch's torch.distributions.Normal.rsample() implements this draw. Using sample() removes this reconstruction path to the encoder, but the decoder can still train, and an analytic KL term can still update encoder parameters.[3]

1import numpy as np
2
3mu = np.array([0.40, -0.20])
4logvar = np.array([-0.70, 0.10])
5epsilon = np.array([0.50, -1.00])
6
7sigma = np.exp(0.5 * logvar)
8z = mu + sigma * epsilon
9
10print("mu:", np.round(mu, 3))
11print("sigma:", np.round(sigma, 3))
12print("epsilon:", np.round(epsilon, 3))
13print("sampled z:", np.round(z, 3))1mu: [ 0.4 -0.2]
2sigma: [0.705 1.051]
3epsilon: [ 0.5 -1. ]
4sampled z: [ 0.752 -1.251]Check the framework contract directly. Predict which draw keeps an autograd path, then verify the derivative through logvar: .
1import torch
2from torch.distributions import Normal
3
4torch.manual_seed(11)
5mu = torch.tensor([0.40, -0.20], requires_grad=True)
6logvar = torch.tensor([-0.70, 0.10], requires_grad=True)
7distribution = Normal(mu, torch.exp(0.5 * logvar))
8ordinary = distribution.sample()
9pathwise = distribution.rsample()
10mu_gradient, logvar_gradient = torch.autograd.grad(
11 pathwise.sum(), (mu, logvar), retain_graph=True
12)
13kl = 0.5 * (mu.square() + torch.expm1(logvar) - logvar).sum()
14kl_mu_gradient = torch.autograd.grad(kl, mu)[0]
15
16print("sample has an autograd path:", ordinary.requires_grad)
17print("rsample has an autograd path:", pathwise.requires_grad)
18print("rsample mean gradient:", mu_gradient.tolist())
19print("logvar derivative matches:", torch.allclose(
20 logvar_gradient, 0.5 * (pathwise.detach() - mu.detach())
21))
22print("analytic KL mean gradient:", [round(x, 3) for x in kl_mu_gradient.tolist()])1sample has an autograd path: False
2rsample has an autograd path: True
3rsample mean gradient: [1.0, 1.0]
4logvar derivative matches: True
5analytic KL mean gradient: [0.4, -0.2]That sampled is what the decoder sees during training. The full path is now:

Generation later skips the encoder. It draws from the prior and runs only the decoder.
The VAE objective
The VAE performs variational inference by optimizing an evidence lower bound (ELBO) on the data log-likelihood.[2] Written as a quantity to maximize:
The first term rewards a decoder that makes the observed crop likely after sampling its latent. The second discourages a posterior that moves too far from the prior. That bound equals only when matches the true posterior . Otherwise the leftover gap is .
In code we minimize the negative ELBO, or a weighted variant:
For , this is the negative ELBO when L_recon estimates the selected negative log-likelihood. To connect that term to the pixel errors you've already calculated, let the decoder output a mean image and assume independent Gaussian pixel noise with fixed variance :
For fixed , the last term doesn't change with the weights. The first is scaled summed squared error, not automatically the mean MSE. A smaller observation variance penalizes the same pixel errors more strongly. Changing the reconstruction reduction without adjusting the KL coefficient changes the optimization problem.
The reconstruction term averages over posterior samples. Our code uses one reparameterized draw per example, giving a Monte Carlo estimate of that expectation. For a fixed model and example with finite variance, averaging more independent draws reduces the estimator's variance in proportion to their count, but requires more decoder work. The Gaussian KL term is evaluated analytically.[2]
This ledger computes both terms for one crop using a fixed hypothetical reconstruction:
1import numpy as np
2
3x = np.array([0.10, 0.90, 0.85, 0.12])
4x_hat = np.array([0.12, 0.84, 0.80, 0.14])
5mu = np.array([0.40, -0.20])
6logvar = np.array([-0.70, 0.10])
7
8reconstruction_mse = np.mean((x - x_hat) ** 2)
9kl = -0.5 * np.sum(1 + logvar - mu**2 - np.exp(logvar))
10
11for beta in [0.0, 0.1, 1.0]:
12 total = reconstruction_mse + beta * kl
13 print(f"beta={beta:.1f}: recon={reconstruction_mse:.4f} KL={kl:.4f} total={total:.4f}")1beta=0.0: recon=0.0017 KL=0.2009 total=0.0017
2beta=0.1: recon=0.0017 KL=0.2009 total=0.0218
3beta=1.0: recon=0.0017 KL=0.2009 total=0.2026Higgins et al. studied -VAE, increasing the KL weight to restrict latent information and encourage disentangled representations.[4] A factorized prior doesn't guarantee that learned coordinates represent independent real-world factors. Unsupervised disentanglement requires assumptions about the model and data.[5] Measure reconstruction, latent usefulness, and prior samples on your task; a higher weight isn't automatically better.
1reconstruction_losses = [0.004, 0.010, 0.028]
2kl_losses = [1.20, 0.42, 0.08]
3betas = [0.01, 0.10, 1.00]
4
5for beta, recon, kl in zip(betas, reconstruction_losses, kl_losses):
6 total = recon + beta * kl
7 print(f"beta={beta:>4.2f}: reconstruction={recon:.3f}, KL={kl:.2f}, objective={total:.3f}")1beta=0.01: reconstruction=0.004, KL=1.20, objective=0.016
2beta=0.10: reconstruction=0.010, KL=0.42, objective=0.052
3beta=1.00: reconstruction=0.028, KL=0.08, objective=0.108These rows are an accounting exercise, not a reported experiment. Each row uses a different objective, so ranking the totals doesn't identify the best model. Log reconstruction and KL separately and inspect samples at each setting.
Train a tiny VAE
Now train a two-dimensional VAE on two families of synthetic crops: one with a bright center highlight and one with a bright lower strip. A constant image can't fit both well. After training, compare reconstruction from each posterior mean with reconstruction from the other family's means. Swapping codes should hurt if they carry pattern information. Finally, decode three prior draws and inspect their center and lower brightness rather than just their tensor shape.
The decoder returns Gaussian mean images in this MSE-based example. We display those means, without adding observation noise. The loss below is mean pixel error plus a chosen KL weight; it doesn't assume that this weight is the standard ELBO coefficient.
1import torch
2from torch import nn
3import torch.nn.functional as F
4
5torch.manual_seed(8)
6center_mark = torch.tensor([
7 0.05, 0.10, 0.08, 0.10, 0.06,
8 0.12, 0.88, 0.86, 0.13, 0.08,
9 0.10, 0.91, 0.90, 0.12, 0.07,
10 0.06, 0.15, 0.14, 0.10, 0.06,
11 0.05, 0.08, 0.08, 0.05, 0.04,
12])
13lower_mark = torch.tensor([
14 0.04, 0.07, 0.06, 0.08, 0.05,
15 0.10, 0.15, 0.12, 0.10, 0.06,
16 0.08, 0.18, 0.16, 0.09, 0.06,
17 0.06, 0.84, 0.88, 0.11, 0.07,
18 0.04, 0.07, 0.08, 0.05, 0.04,
19])
20patches = torch.cat([
21 torch.clamp(center_mark + 0.02 * torch.randn(32, 25), 0.0, 1.0),
22 torch.clamp(lower_mark + 0.02 * torch.randn(32, 25), 0.0, 1.0),
23])
24
25class TinyVAE(nn.Module):
26 def __init__(self):
27 super().__init__()
28 self.encoder = nn.Sequential(nn.Linear(25, 12), nn.Tanh())
29 self.mu = nn.Linear(12, 2)
30 self.logvar = nn.Linear(12, 2)
31 self.decoder = nn.Sequential(
32 nn.Linear(2, 12), nn.Tanh(), nn.Linear(12, 25), nn.Sigmoid()
33 )
34
35 def forward(self, x):
36 hidden = self.encoder(x)
37 mu = self.mu(hidden)
38 logvar = self.logvar(hidden)
39 z = mu + torch.exp(0.5 * logvar) * torch.randn_like(mu)
40 return self.decoder(z), mu, logvar
41
42model = TinyVAE()
43optimizer = torch.optim.Adam(model.parameters(), lr=0.03)
44beta = 0.02
45
46def loss_terms(x_hat, x, mu, logvar):
47 recon = F.mse_loss(x_hat, x)
48 kl = -0.5 * (1 + logvar - mu.pow(2) - logvar.exp()).sum(1).mean()
49 return recon + beta * kl, recon, kl
50
51for _ in range(400):
52 x_hat, mu, logvar = model(patches)
53 loss, recon, kl = loss_terms(x_hat, patches, mu, logvar)
54 optimizer.zero_grad()
55 loss.backward()
56 optimizer.step()
57
58torch.manual_seed(99)
59with torch.no_grad():
60 x_hat, mu, logvar = model(patches)
61 loss, recon, kl = loss_terms(x_hat, patches, mu, logvar)
62 generated = model.decoder(torch.randn(3, 2))
63 mean_mse = F.mse_loss(model.decoder(mu), patches)
64 swapped_mse = F.mse_loss(model.decoder(mu.roll(32, dims=0)), patches)
65 kl_per_dim = -0.5 * (1 + logvar - mu.square() - logvar.exp()).mean(0)
66
67print(f"reconstruction={recon.item():.4f} KL={kl.item():.4f} total={loss.item():.4f}")
68print("generated batch shape:", tuple(generated.shape))
69print(f"posterior-mean MSE={mean_mse.item():.4f}; swapped-code MSE={swapped_mse.item():.4f}")
70print("KL per dimension:", [round(value, 4) for value in kl_per_dim.tolist()])
71for index, crop in enumerate(generated):
72 center = crop[[6, 7, 11, 12]].mean().item()
73 lower = crop[[16, 17]].mean().item()
74 print(f"prior draw {index}: center={center:.3f}, lower={lower:.3f}")1reconstruction=0.0075 KL=0.7382 total=0.0222
2generated batch shape: (3, 25)
3posterior-mean MSE=0.0005; swapped-code MSE=0.1276
4KL per dimension: [0.0005, 0.7378]
5prior draw 0: center=0.142, lower=0.875
6prior draw 1: center=0.877, lower=0.150
7prior draw 2: center=0.800, lower=0.219Each object in the program now has a job: mu and logvar define each approximate posterior, torch.randn_like supplies , decoder consumes sampled , and the logged losses reveal the compromise.
The VAE's shape contract is small enough to write down. Here, B is batch size:
| Tensor | Shape in this program | Role |
|---|---|---|
Input batch x | [B, 25] | B flattened 5 by 5 crops |
| Encoder hidden state | [B, 12] | Shared features before the two heads |
mu, logvar, and sampled z | [B, 2] | One two-dimensional posterior per crop |
Reconstruction x_hat | [B, 25] | One flattened crop per sampled latent |
The reductions matter too. F.mse_loss averages pixel errors over the batch, while the KL term sums its two latent dimensions before averaging over the batch. That makes beta = 0.02 a scale choice for this setup, not a universal VAE setting. If you change latent width or loss reduction, retune it and keep logging both terms.
Swapping the families' codes raises MSE from 0.0005 to 0.1276, showing that the decoder uses input-dependent information. Prior draw 0 has a bright lower strip; draw 1 has a bright center; draw 2 is less distinct. Three draws can't establish distributional quality, and the reconstruction checks use training crops. Add held-out patterns and inspect full decoded grids before making a generalization claim.
Watch for posterior collapse
Low KL can look like success because each posterior matches the prior. But if reconstruction stays acceptable while changing barely changes the output, what happened? The decoder found a shortcut and stopped using its latent input. In information-theoretic terms, the mutual information between observations and latent codes can fall toward zero: knowing the sampled code no longer tells you much about which crop produced it. This failure is called posterior collapse. Van den Oord et al. describe it as a common issue for VAEs paired with high-capacity autoregressive decoders.[6]
An expressive autoregressive decoder can model data through previously observed pixels or tokens while ignoring the global latent. Training can settle near , with little input-specific information in the code.[7] That decoder can still generate varied, plausible outputs. An unconditional deterministic mean-image decoder optimized with MSE instead favors the dataset's mean image when it ignores . Averaging is a consequence of that decoder and loss, not a definition of collapse.
There are two pressures hidden in the average KL. Let be the mixture of encoder posteriors over the data and let measure information under that joint data/encoder distribution. Then:[8]
Matching the aggregate codes to the prior and making every crop's posterior equal to the prior are different goals. If every posterior equals the prior, : a sampled code cannot tell us which crop it came from. A large average KL, however, may reflect prior mismatch rather than useful information.
The training code already reports KL per latent dimension. One near-zero dimension can simply be unused capacity when the other distinguishes the two patterns. This illustrative monitor raises a warning only when KL is tiny in every dimension; it doesn't diagnose collapse by itself:
1import numpy as np
2
3logged_kl_by_dimension = {
4 "using latent": np.array([0.31, 0.18]),
5 "suspected collapse": np.array([0.0003, 0.0001]),
6}
7threshold = 0.001
8
9for run_name, kl_dims in logged_kl_by_dimension.items():
10 low_kl_warning = bool(np.all(kl_dims < threshold))
11 print(f"{run_name:18} total_KL={kl_dims.sum():.4f} low_KL_warning={low_kl_warning}")1using latent total_KL=0.4900 low_KL_warning=False
2suspected collapse total_KL=0.0004 low_KL_warning=TrueLow KL alone isn't a proof of failure. Pair it with reconstruction quality, generated samples, and checks that predictions change when changes. Still, a rapid collapse toward zero is worth investigating.
Three interventions worth testing are:
- KL annealing: Start with weak KL pressure and increase it toward the intended final weight. Bowman et al. anneal from zero to one for their likelihood-scaled objective.[7] Our mean-MSE example uses a different scale, so one isn't a universal endpoint. Zero KL weight also doesn't guarantee useful latents, as the nearly constant-crop experiment showed.
- Free bits: The original formulation sums over latent groups, with each group's KL averaged over the minibatch before applying the floor.[9] Below , that KL contribution has zero gradient; reconstruction can reward additional information without extra KL cost. This doesn't force a minimum mutual information. With natural logarithms, the threshold is measured in nats despite the name.
- Decoder rebalancing: Reduce decoder capacity or weaken autoregressive conditioning. Bowman's word dropout removes some teacher-forced context so latent information becomes more useful.[7] It doesn't establish that the decoder must use , and degrading a decoder can hurt quality.
Inspect interpolation rather than assuming it is meaningful. A decoder built from continuous layers changes continuously with in a deterministic autoencoder too; that mathematical property doesn't guarantee plausible images. The path can cross regions poorly represented during training in either model. Gaussian posterior support alone doesn't tell us where most probability mass lies or where decoding works well.
An expressive autoregressive VAE has KL near zero but generates varied, plausible images. Does that rule out posterior collapse?
Answer
No. The decoder can model a rich distribution using previous pixels while ignoring the global latent. Test whether the encoder's code carries input-specific information and whether changing that code changes the conditional output distribution. Sample variety alone doesn't establish latent use.
VQ-VAE: replace continuous samples with codes
A Vector Quantised Variational Autoencoder (VQ-VAE) takes a different route. Its encoder output is mapped to the nearest vector in a learned codebook, so its latent representation is discrete rather than Gaussian. The original VQ-VAE paper also learns a prior over those discrete codes.[6]
For one encoder vector, quantization is a nearest-neighbor lookup:
1import numpy as np
2
3encoder_output = np.array([0.72, -0.10])
4codebook = np.array([
5 [-0.80, 0.20],
6 [0.65, -0.05],
7 [0.10, 0.90],
8])
9
10squared_distances = ((codebook - encoder_output) ** 2).sum(axis=1)
11index = int(squared_distances.argmin())
12quantized = codebook[index]
13
14print("squared distances:", np.round(squared_distances, 3))
15print("selected code ID:", index)
16print("decoder receives:", quantized)1squared distances: [2.4 0.007 1.384]
2selected code ID: 1
3decoder receives: [ 0.65 -0.05]Nearest-code assignments are piecewise constant, so their ordinary derivatives don't supply a useful encoder training signal. VQ-VAE uses a straight-through gradient estimator for the encoder path, plus a codebook loss and a commitment loss so encoder outputs stay close to useful entries:[6]
Here stops gradients, is the selected codebook vector, and this is the VQ-VAE commitment weight, not the -VAE KL weight. In the original straight-through setup, reconstruction trains the decoder and encoder, but supplies no codebook gradient. The codebook loss updates selected entries; commitment updates the encoder. A moving-average codebook update is another option.[6] The important distinction is the interface:
The nearest code doesn't change during this backward pass. Which tensors receive gradients from reconstruction alone, and which need the auxiliary losses? This two-coordinate example makes the routing visible:
1import torch
2
3encoded = torch.tensor([0.72, -0.10], requires_grad=True)
4codebook = torch.tensor([
5 [-0.80, 0.20], [0.65, -0.05], [0.10, 0.90]
6], requires_grad=True)
7index = ((codebook - encoded) ** 2).sum(1).argmin()
8quantized = codebook[index]
9straight_through = encoded + (quantized - encoded).detach()
10target = torch.tensor([0.50, 0.10])
11reconstruction = 0.5 * (straight_through - target).square().sum()
12encoder_recon, codebook_recon = torch.autograd.grad(
13 reconstruction, (encoded, codebook), retain_graph=True, allow_unused=True
14)
15codebook_loss = (encoded.detach() - quantized).square().sum()
16commitment = 0.25 * (encoded - quantized.detach()).square().sum()
17(reconstruction + codebook_loss + commitment).backward()
18
19rounded = lambda vector: [round(x, 3) for x in vector.tolist()]
20print("forward code:", rounded(straight_through.detach()))
21print("reconstruction encoder gradient:", rounded(encoder_recon))
22print("reconstruction codebook gradient:", codebook_recon)
23print("total encoder gradient:", rounded(encoded.grad))
24print("selected code gradient:", rounded(codebook.grad[index]))
25assert torch.count_nonzero(codebook.grad[[0, 2]]).item() == 01forward code: [0.65, -0.05]
2reconstruction encoder gradient: [0.15, -0.15]
3reconstruction codebook gradient: None
4total encoder gradient: [0.185, -0.175]
5selected code gradient: [-0.14, 0.1]| Bottleneck | Decoder receives | How a new latent is obtained |
|---|---|---|
| Autoencoder | Deterministic vector | Encode an input |
| VAE | Continuous sampled vector | Sample from a prior after training |
| VQ-VAE | Codebook vector selected by an ID | Sample code IDs from a learned prior after training |
A hand-picked VQ-VAE code ID can be decoded, but it isn't automatically a plausible sample. The learned prior supplies the combinations of code IDs that the generative model has been trained to use.
Continuous VAEs and discrete VQ-VAEs can both turn a large input into a smaller latent another model can consume. A later generator can learn the distribution of those latents instead of modeling every pixel directly.
Latent models reduce the expensive space
The same bottleneck contract scales beyond this tiny crop. A later generator can work in latent space, then hand its final representation to the decoder that turns it back into pixels. Rombach et al. train diffusion models in the latent space of pretrained autoencoders rather than applying all denoising operations directly to pixels.[10] They write the spatial downsampling factor as and evaluate the efficiency and quality trade-off across values from 2 to 32, including and . The autoencoder's job is precise: compress before the expensive generation steps, then decode a final latent back to pixels.
Their compression stage uses perceptual and adversarial losses, with weak KL or vector-quantization regularization. A separate diffusion model learns the latent distribution.[10] Sampling Gaussian noise and decoding it directly isn't equivalent to running that generator: diffusion must first transform noise into a latent that fits the learned distribution.
The size reduction itself is simple arithmetic. A downsampling factor of 8 reduces each spatial axis by 8, so it reduces spatial positions by 64. On a 512 by 512 grid that is 262,144 positions becoming 4,096:
1height, width = 512, 512
2downsample = 8
3latent_height = height // downsample
4latent_width = width // downsample
5
6pixel_positions = height * width
7latent_positions = latent_height * latent_width
8
9print("pixel grid:", (height, width), "positions:", pixel_positions)
10print("latent grid:", (latent_height, latent_width), "positions:", latent_positions)
11print("spatial reduction:", pixel_positions // latent_positions, "x")1pixel grid: (512, 512) positions: 262144
2latent grid: (64, 64) positions: 4096
3spatial reduction: 64 xChannels and decoder quality still matter. For a hypothetical RGB image and four-channel latent at , scalar values drop from to , a 48-fold reduction, rather than 64. Neither ratio predicts runtime or compressed-file size by itself. Rombach et al. find that aggressive compression hurts their downstream results; detail discarded by the encoder can't be uniquely recovered from its code. A generator may synthesize plausible detail without reproducing the original pixels.
Debug the latent interface
Start by logging reconstruction and KL losses separately. A low total can hide a reconstruction term that dominates, or a KL term that has collapsed toward zero.
If sampling or KL returns NaN or Inf, check the scale path first: sigma = exp(0.5 * logvar), not exp(logvar). Inspect the range of logvar before adding an explicit clamp, since a clamp can hide unstable optimization.
Keep the two distributions distinct. The encoder supplies for an observed crop; generation samples from the prior . If decoded outputs barely change when changes, the decoder may be ignoring its latent input even when reconstruction looks fine.
For a VQ-VAE, log code usage as well. If only a few code IDs are selected, the codebook's capacity is going unused and the discrete bottleneck needs a closer look.
Test what the code preserves
What changes when an autoencoder becomes a VAE?
Answer
The encoder no longer emits only one deterministic latent vector. It predicts parameters of a posterior distribution, the model samples a latent through reparameterization, and training adds a KL term that relates that posterior to a prior used for generation.
Why does the VAE use z = mu + sigma * epsilon?
Answer
The random draw is moved into epsilon, which is independent of the encoder parameters. For a fixed sampled epsilon, z remains a differentiable function of mu and sigma, so reconstruction gradients reach the encoder.
Why should you log KL separately from reconstruction loss?
Answer
A small total can hide failure. Very weak KL can leave poor prior sampling, while KL collapsing near zero can indicate that the decoder ignores z. The two terms explain different behaviors.
The trained autoencoder and a constant mean-crop predictor have nearly equal MSE. Has the two-number latent learned a useful representation?
Answer
That result doesn't establish a useful representation. The dataset may be so uniform that the decoder can reproduce the shared pattern while ignoring most input variation. Compare against distinct held-out patterns and test whether changing or swapping codes changes reconstruction error.