primer.ml.generative.autoencoders
Autoencoders and VAEs
Run: python -m primer.ml.generative.autoencoders
This lesson builds on the two-layer network and training loop of
primer.ml.neural_net and the Adam optimizer of primer.ml.optimizers;
primer.notation explains every symbol from zero.
Level 1: The practitioner's guide
In one sentence. An autoencoder is a pair of networks trained to squeeze data through a narrow code and rebuild it, which makes it a learned lossy compressor; a variational autoencoder (VAE) also shapes the code space so that a random code decodes to something sensible, which makes it a generator and, far more often today, the compressed space that bigger generators (diffusion models, transformers) work inside.
When you need it. Three tells. You have unlabelled data and want a compact, meaningful representation of each item (a fingerprint for search, a small input for another model, a way to flag the items that don't fit): that is a plain autoencoder. You want to generate or edit new items and need a smooth space where nearby codes mean similar outputs: that is a VAE. You are building or running an image, audio or video generator: there is almost certainly an autoencoder at both ends of it, and its settings (downsampling factor, latent channels, scale factor, precision) are yours to get right. You don't need one for compression that must be exact (a zip file is lossless; an autoencoder never is), and you rarely train one for images or audio yourself any more: pretrained ones are downloadable, and their quality took a great deal of data to reach. The number that shows the naive path failing: in this lesson a plain autoencoder rebuilds 8 × 8 pen strokes from two numbers with an error of 0.07 (99% of the picture kept, against 50% for PCA), yet decoding random codes drawn from the standard bell curve gives junk 76% of the time. Compression is not generation.
Your options. From the cheapest to the most capable:
| Option | What it does | What it guarantees | What it costs | Where it lives |
|---|---|---|---|---|
| PCA | Fits a flat sheet through the data; an item's code is where it lands on the sheet | The best any flat code can do under squared error; no training loop | One matrix decomposition; poor rebuilds of curved data (50% kept here) | A library call |
| Plain autoencoder | A bent encoder and decoder trained only to rebuild the input | Far better rebuilds on curved data (99% kept here); a code that flags anomalies and can denoise | A training run; a code space with holes, so no generation | Your training loop |
| VAE | The encoder emits a fuzzy region per item and pays rent for straying from the standard bell curve | Random codes decode sensibly (36% junk here, against 76%); smooth interpolation | Blurrier rebuilds; a β to tune; posterior collapse if you overdo it | Your training loop, or a pretrained one |
| Latent autoencoder for diffusion | A VAE with a tiny KL weight; the diffusion model generates inside its code space | 48 times fewer numbers for a 512 × 512 image, so training and sampling become affordable | A ceiling on detail set by the decoder; scaling, precision and memory settings to respect | Downloaded with the diffusion model |
| VQ-VAE and VQGAN | Snaps each code vector to the nearest entry of a learned codebook, so an item becomes a grid of token ids | Tokens a transformer reads and writes like words; sharper output when an adversarial loss is added (VQGAN) | A codebook to keep in use; a discrete space with no straight-line interpolation | A pretrained tokenizer |
| Neural audio codec | The same recipe for sound: encoder, residual quantizer, decoder, reconstruction plus adversarial losses | Speech and music at 3 to 18 kbit/s, streamable in real time | A model at both ends of the wire | SoundStream, EnCodec |
How to choose. Start from what you want the code for.
- Compact features for search, clustering or a downstream model, and no generation: try PCA first (Level 2 shows it is exactly the autoencoder with no bends). Train an autoencoder when the data is curved and PCA's rebuilds are poor, as they are here.
- Flagging oddities (fraud, a failing machine): a plain autoencoder trained on normal data, with an alarm on rebuild error.
- Generating or editing new items in a small domain, or sliders that mean something: a VAE, with β chosen by looking at samples, not at the loss.
- Generating images, audio or video at real resolution: don't generate with the VAE. Use it as the compressor and let diffusion or a transformer do the generating, and download the autoencoder that generator was trained with, because the pair is matched.
- Feeding images or audio into a language-model-style transformer: a VQ tokenizer or a neural codec.
- Whatever you pick, look at the rebuilds before anything else. The decoder's rebuild quality is the ceiling of every generator that works in its space; no generator can produce detail the decoder cannot paint.
What it costs. Training cost is dominated by data: the autoencoder must
see enough of the domain to rebuild it. At run time it is cheap: one
encoder pass on the way in and one decoder pass on the way out, against the
many passes of the generator between them. That is the economics of latent
diffusion. Rombach et al. tried downsampling factors
from 1 (raw pixels) to 32 and found factors 4 and 8 the sweet spot; after
two million training steps the pixel-space model trailed the factor-8 model
by 38 FID points, and on inpainting the latent models ran at least 2.7 times
faster. Memory: at high resolution the decoder's activations fill the GPU,
which is why the diffusers AutoencoderKL offers tiled encoding and
decoding (constant memory, at the risk of faint tile seams) and runs the
SDXL autoencoder in float32 by default. Quality has a measured ceiling: the
Stable Diffusion 3 paper reports that widening the latent from 4 to 8 to 16
channels drops reconstruction FID from 2.41 to 1.56 to 1.06 and raises PSNR
from 25.12 to 26.40 to 28.62, which is why newer models carry 16-channel
latents at the price of a harder generation task. β is a dial with a bottom:
in this lesson, samples land nearest real strokes at β = 0.3 (median
distance 0.56) and further at both β = 0.01 (2.40) and β = 3 (4.73). Tokens
cost context: DALL-E's discrete VAE turns a 256 × 256 image into 32 × 32 =
1,024 tokens from a codebook of 8,192, cutting the transformer's context
192-fold.
What breaks.
- Holes. A plain autoencoder's codes land wherever training put them (from −24 to 14 here), with empty fields between; 44% of random codes drawn even from inside that range decode to junk. If you need to sample, you need the KL rent: use a VAE.
- Blur. Under squared error the best guess for an uncertain pixel is the average, so a VAE whose regions overlap paints averages. Lower β for sharper rebuilds, or add an adversarial loss as VQGAN does, or stop asking the VAE to generate and let a stronger model do it in its space.
- Posterior collapse. Raise β too far and every region becomes the standard bell curve, the code carries nothing, and the decoder emits the same average picture for every code (β = 3 here: KL exactly 0, rebuild error 6.08, one grey smudge). Watch the KL term; zero is a symptom.
- Forgetting the latent scale. Diffusion libraries multiply latents by a scaling factor (0.18215 for Stable Diffusion's autoencoder) so they have unit variance for the generator, and divide it back out before decoding. Skip either step and the generator sees data it never trained on.
- Precision. The Stable Diffusion autoencoders overflow in float16 at high resolution; run them in float32 or use a checkpoint fine-tuned for half precision.
- Dead codebook entries. In a VQ model, entries nothing maps to waste the vocabulary. DALL-E raised its KL weight to 6.6 to promote codebook usage; if your tokens cluster on a few ids, that is the dial.
- A mismatched pair. Latents from one autoencoder decoded by another are junk. Keep the encoder, the generator and the decoder that were trained together.
In the wild. Stable Diffusion's KL-regularised autoencoder (8 times
downsampling, 4 latent channels) ships with every Stable Diffusion model
and is AutoencoderKL in Hugging Face diffusers; Stable Diffusion 3 moved
to 16 channels. VQGAN (Esser, Rombach and Ommer) adds an adversarial loss
to a VQ-VAE and puts a transformer over its tokens; the original VQ-VAE
compressed 128 × 128 images to a 32 × 32 grid over a codebook of 512, about
42.6 times fewer bits, and generated with a PixelCNN over the grid; DALL-E's
discrete VAE with 8,192 codes fed its text-to-image transformer. For sound,
SoundStream (a convolutional encoder and decoder around a residual vector
quantizer, 3 to 18 kbit/s, and at 3 kbit/s preferred over Opus at 12) and
EnCodec (a streaming encoder-decoder with a multiscale spectrogram
adversary, at 24 kHz mono and 48 kHz stereo) are the same recipe at audio's
scale. Every paper is linked at the end of the lesson.
Go deeper. Level 2 builds both halves by hand on 8 × 8 pen strokes, shows where a plain autoencoder's holes come from, then adds the VAE's two pieces (the reparameterization trick and the KL rent) with numbers you can check, and sweeps β to watch holes give way to blur and then collapse. If you only needed to choose, you are done.
Level 2: How it works, from scratch
You phone a friend and describe a picture so they can draw it, but you are allowed to say only two numbers. You would agree on a system first: the first number says which way the line runs, the second says how far it sits from the middle. You squeeze the picture into two numbers, send them, and your friend rebuilds it.
An autoencoder is a neural network that invents that system by itself. It has two halves. The encoder squeezes a picture into a few numbers, called the code. The decoder rebuilds the picture from the code. Nobody tells it what the numbers should mean. The only instruction is "the rebuild must match the original", and to obey it through such a narrow gap the network has to discover what really varies in the data.
Then comes a second question. If your friend can draw from any two numbers,
can you invent a new picture by making up two numbers? With a plain
autoencoder, usually not: most made-up codes were never used for anything,
and the drawing comes out as a smudge. A variational autoencoder (VAE) is
trained so that a made-up code drawn from a known range draws something
sensible. That step, from compressing to generating, is what this lesson is
about, and it opens this part of the primer: GANs
(primer.ml.generative.gans) and diffusion (primer.ml.generative.diffusion)
are other answers to the same question.
The toy data: pictures of pen strokes
Tiny example. Every picture in this lesson is 8 × 8 = 64 pixels holding one soft pen stroke. The stroke runs in one of four directions (horizontal, vertical, diagonal down, diagonal up) and is shifted up to 2.5 pixels from the centre. The ink fades like a bell curve either side of the line: a pixel on the line has brightness 1.0, a pixel one pixel away about 0.25, a pixel two away almost nothing.
So each picture is 64 numbers, but only two facts change from picture to picture: the direction (one of four) and the offset (any amount). A good 2-number code has to rediscover both from the pixels alone. The 200 pictures are the top row of the first figure below.
In code: make_strokes draws the pictures and returns their hidden directions and offsets; STROKE_KINDS names the four directions.
Squeeze and rebuild: the autoencoder
Everyday picture. A zip file squeezes a document and gets it back exactly. An autoencoder is a lossy, learned zip: it keeps what matters most for this kind of data and lets the rest go.
Tiny example: a one-number code you can check by hand. Take points that lie on the line y = 2x, such as (1, 2) and (2, 4). Each takes two numbers to write down, but one number is enough: how far along the line the point is.
- Encoder: code = u · x, where u = (1, 2)/√5 = (0.447, 0.894) is the
line's direction scaled to length 1, and "·" is the dot product (multiply
matching entries and add; see
primer.notation). - Decoder: rebuild = code × u.
For (1, 2): code = 0.447 + 1.789 = 2.236 (which is √5, its distance from the origin), and rebuild = 2.236 × (0.447, 0.894) = (1, 2). Nothing lost.
For (2, 3), which is off the line: code = (2 + 6)/√5 = 3.578, and rebuild = 3.578 × (0.447, 0.894) = (1.6, 3.2), the closest point on the line. The part that pointed away from the line, (0.4, −0.2), is gone. That loss is exactly what training measures and shrinks:
Level 3: the formula and its symbols
$$ L_{\text{rec}} = \frac{1}{n} \sum_{i=1}^{n} \bigl\lVert x_i - g\bigl(f(x_i)\bigr) \bigr\rVert^2 $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $x_i$ | the $i$-th example: a point, or a picture's 64 pixels | $x = (2, 3)$ |
| $f$ | the encoder: squeezes an example into a code | $f(x) = u \cdot x$ |
| $f(x_i)$ | the code for example $i$, often written $z_i$ | 3.578 |
| $g$ | the decoder: rebuilds an example from a code | $g(z) = z\,u$ |
| $g(f(x_i))$ | the rebuild, often written $\hat{x}_i$ ("x hat") | (1.6, 3.2) |
| $x_i - \hat{x}_i$ | what the rebuild got wrong, entry by entry | (0.4, −0.2) |
| $\lVert v \rVert^2$ | squared length of $v$: square every entry and add | $0.4^2 + 0.2^2 = 0.2$ |
| $\sum_{i=1}^{n}$ | add up over every example | |
| $n$ | how many examples | 2: (1, 2) and (2, 3) |
| $L_{\text{rec}}$ | the reconstruction loss: average squared error of the rebuilds | 0.1 |
In words: "squeeze each example, rebuild it, measure the squared distance between the rebuild and the original, and average over all the examples."
With the numbers: (1, 2) rebuilds perfectly, error 0. (2, 3) rebuilds as (1.6, 3.2), error 0.4² + (−0.2)² = 0.16 + 0.04 = 0.2. Averaged over the two points, $L_{\text{rec}}$ = (0 + 0.2)/2 = 0.1.
Level 3: in Python
In Python:
import math
# u: the line's direction, scaled to length 1
u = [1 / math.sqrt(5), 2 / math.sqrt(5)]
def f(x):
return sum(u_m * x_m for u_m, x_m in zip(u, x))
def g(z):
return [z * u_m for u_m in u]
def error(x):
return sum((a - b) ** 2 for a, b in zip(x, g(f(x))))
# f(x): the encoder squeezes (2, 3) to one number
round(f([2, 3]), 3) # → 3.578
# g(f(x)): the decoder rebuilds two numbers from it
[round(v, 2) for v in g(f([2, 3]))] # → [1.6, 3.2]
# ‖x − x̂‖² for each point: (1, 2) is on the line, (2, 3) is not
[round(error(x), 2) for x in ([1, 2], [2, 3])] # → [0.0, 0.2]
# L_rec: the average over n = 2 points
round((error([1, 2]) + error([2, 3])) / 2, 2) # → 0.1
The real network has the same two halves, only bent. The encoder is a
two-layer network like the one built in primer.ml.neural_net: 64 pixels
into 32 hidden numbers (through tanh, which lets it bend), then out to a code
of 2 numbers. The decoder mirrors it: 2 numbers into 32 hidden, then out to
64 pixels squashed between 0 and 1 by a sigmoid so they are valid
brightnesses. Training is the loop from primer.ml.neural_net with the Adam
optimizer from primer.ml.optimizers: rebuild every picture, measure
$L_{\text{rec}}$, send the gradient back through both halves, adjust, 1,500
times.
flowchart LR X["picture x<br/>64 pixels"] --> E["encoder f<br/>64 → 32 → 2"] E --> Z["code z<br/>2 numbers"] Z --> D["decoder g<br/>2 → 32 → 64"] D --> XH["rebuild x̂<br/>64 pixels"] X --> L["loss<br/>squared error between x and x̂"] XH --> L L -. "gradients flow back<br/>through decoder, then encoder" .-> E
Reading it: follow the picture left to right. It is 64 numbers wide at both ends and only 2 numbers wide in the middle: that narrow middle is the bottleneck, and it is the whole point. Without it the network could copy the pixels straight through and learn nothing. The loss box compares the two ends, and the dotted arrow is backpropagation carrying the blame back through the decoder and on into the encoder, so both halves learn together. The encoder never sees a target code; it learns whatever code the decoder finds most useful.
Reading it: the top row is eight real strokes, two of each direction. The middle row is what the autoencoder rebuilds from just 2 numbers per picture: nearly identical. The average error is about 0.07, against about 8.8 units of squared ink in a whole stroke, so it keeps over 99% of the picture. The bottom row squeezes the same pictures to 2 numbers with PCA and rebuilds them: grey smudges, keeping only about half. Same budget of two numbers, very different results. The next section explains why.
In code: worked_example_line is the hand example above; Autoencoder holds both halves and its hand-written backward pass in Autoencoder.loss_and_gradients; train runs Adam; trained_autoencoder is the lesson's trained network; reconstruction_error measures the rebuilds.
PCA is the straight-line special case
Everyday picture. PCA (principal component analysis, built in
primer.ml.embeddings.clustering) fits a flat sheet through the data and
records where each point lands on the sheet. That is an autoencoder with no
bends: the encoder is one matrix multiply, and so is the decoder.
Tiny example. Scatter 100 points near the line y = 2x and train exactly that no-bend autoencoder, with a 1-number code, by plain gradient descent. The direction it learns is (0.447, 0.894), the line's own direction, and its rebuild error matches PCA's with one component. Baldi and Hornik proved in 1989 that this always happens: a linear autoencoder trained on squared error lands on the same flat sheet as PCA.
So why did PCA smudge the strokes? Because a stroke sliding across the image does not travel in a straight line through pixel space. Take a horizontal stroke at the top and another at the bottom. Their average, the point halfway along the straight line between them, is two faint lines, not one stroke in the middle. The real strokes lie on a curved surface in the 64-dimensional space of pictures, and a flat sheet can only cut through it. The autoencoder's tanh layers let its "sheet" bend to follow the curve.
Why it matters in practice. Before generation, autoencoders earned their keep as learned compressors. They compress, they denoise (train on noisy inputs, ask for clean outputs), and they detect anomalies: an input the network rebuilds badly is unlike anything it trained on, which is a standard way to flag fraud or a failing machine.
Reading it: each dot is one picture, placed at its 2-number code and coloured by its stroke direction, with bigger dots for larger offsets. The codes form strands: pictures of one direction line up along a curve (now and then broken into pieces), and sliding along a strand slides the stroke. The network rediscovered both hidden facts without being told either. Now look at the axes: the codes run from about −24 to 14. Nothing asked for that range; it is an accident of training. The dashed circle near 0 is where a "random" code from the bell curve would usually land, and it catches almost none of the strands. Hold on to that for the next section.
In code: train_linear_autoencoder is the no-bend autoencoder; pca_reconstruction rebuilds from the top principal components.
Why a plain autoencoder cannot generate: holes
Everyday picture. Imagine a town where houses were built only along four winding roads. Pick a random spot inside the town limits and you will most likely land in a field. Ask the decoder to draw the picture that "lives" in that field, and it answers anyway, with whatever its weights happen to produce, because nothing in training ever asked it about that spot.
Tiny example. The obvious way to make up a random code is to draw each number from the standard normal distribution, the bell curve centred on 0 with a spread (standard deviation) of 1: most draws fall between −1 and 1, almost all between −3 and 3. Written $\mathcal{N}(0, I)$ for several numbers at once, it means "draw each number from that bell curve, independently".
Draw the code z = (1.2, −0.9) and decode it with the plain autoencoder. The result is a bright smear with about three times the ink of any real stroke; its squared distance to the nearest real stroke is about 7.5. Do that 500 times and the median distance is about 3, while three quarters of the draws land further than 1.0 from every real stroke. That 1.0 is this lesson's line for junk: about a ninth of a whole stroke's squared ink.
To be fair to the autoencoder, draw instead from the box that holds its own codes, from −24 to 1.6 across and −17 to 14 up. Still about 44% of those codes decode to junk. They fall in the holes between the strands.
Reading it: on the left, black dots are the codes of real strokes and the dashed rectangle is the box around them. Every blue circle is a random code that decoded to something close to a real stroke; every red cross decoded to junk. The red crosses sit in the open spaces between strands, the blue circles near them. On the right are the decoded pictures themselves: the 12 worst (top two rows) are black blobs no pen would draw, and the 12 best (bottom) are clean strokes, because those random codes happened to land on a strand.
Why it matters in practice. A compressor is not a generator. To sample new data you need a code space with a known shape that you can draw from, no holes inside that shape, and smoothness, so nearby codes decode to similar pictures. A plain autoencoder promises none of these, because its loss only ever looks at codes of real pictures.
In code: generate decodes codes drawn from the standard normal; codes_in_box draws from a model's own code range; distance_to_data measures how far each image is from the nearest real stroke, and JUNK_DISTANCE is the line between plausible and junk.
The VAE: encode to a fuzzy region, not a point
Everyday picture. Instead of pinning each picture to one exact spot on the map, the encoder draws a small fuzzy circle: "somewhere around here". During training the decoder is handed a random point from inside that circle, so it must draw the right picture from anywhere nearby. A whole neighbourhood now decodes sensibly, not just one pin. A second rule stops the encoder from cheating: every circle pays rent, more the further it sits from the middle of the map and more if it shrinks towards a pin. Circles crowd towards the centre and overlap, and the fields between the roads fill in.
Tiny example. One picture, one code number. The encoder says μ = 0.5 and σ = 0.5: "about 0.5, give or take 0.5". This step, a bell-curve roll gives ε = 1.2, so the decoder is handed z = 0.5 + 0.5 × 1.2 = 1.1. Next step the roll is ε = −0.4 and the decoder gets z = 0.5 − 0.2 = 0.3. Both must rebuild the same picture.
flowchart LR X["picture x"] --> E["encoder"] E --> MU["μ: centre of the region"] E --> LV["log σ²: size of the region"] EPS["ε drawn from N(0, I)"] --> Z["z = μ + σ·ε"] MU --> Z LV --> Z Z --> D["decoder"] --> XH["rebuild x̂"] XH --> R["rebuild error<br/>‖x − x̂‖²"] X --> R MU --> KL["KL rent<br/>pulls regions to N(0, I)"] LV --> KL R --> LOSS["loss = rebuild + β · KL"] KL --> LOSS
Reading it: compare it with the plain autoencoder's diagram. The encoder now has two outputs per code number: a centre μ and a size, given as log σ² (the logarithm of the variance, used because it can be any number while σ itself must stay positive). The box z = μ + σ·ε is where the fuzzy region becomes one concrete code: ε is fresh random noise every step. The loss has two parts. The rebuild error, as before, wants each region small and distinct so the decoder knows exactly which picture it came from. The KL rent, fed straight from μ and log σ², wants every region to look like the standard bell curve. Training settles on a compromise between them, and β sets the exchange rate.
In code: VAE is the Autoencoder with a two-headed encoder; VAE.encode_distribution returns μ and log σ², and VAE.encode returns just the centre μ, the best single code for a picture.
The reparameterization trick: moving the dice outside
Everyday picture. To pick a random seat in a row, you could close your eyes and point. If someone then moves the row one seat to the left, you have no idea how your pick would have changed. Or you could roll a die for "how many seats from the middle" and count from wherever the middle is. Now if the row moves one seat left, your seat moves exactly one seat left. Same kind of random seat, but you can say how it responds to the row moving.
Training needs exactly that. Backpropagation asks of every step: "if I nudge
this number, how does the loss change?" (that rate of change is the
gradient, or derivative; see primer.notation). A raw dice roll has
no answer, so the gradient would stop at the sampling step and the encoder
would never learn. The reparameterization trick rolls the dice first, as
an ordinary input ε, and builds the code with arithmetic:
Level 3: the formula and its symbols
$$ z = \mu + \sigma \odot \varepsilon, \qquad \varepsilon \sim \mathcal{N}(0, I), \qquad \sigma = e^{\frac{1}{2} \log \sigma^2} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $\mu$ | the centre of the picture's region, from the encoder | 0.5 |
| $\log \sigma^2$ | the natural logarithm of the region's variance, from the encoder; "log of y" is the power you raise $e$ to in order to get y | $\log 0.25 = -1.386$ |
| $e$ | Euler's number, ≈ 2.718 | |
| $\sigma$ | the spread of the region (standard deviation); $e^{\frac{1}{2}\log\sigma^2}$ undoes the log and the square | 0.5 |
| $\varepsilon$ | "epsilon": random noise, one number per code number | 1.2 |
| $\sim$ | "is drawn from" | |
| $\mathcal{N}(0, I)$ | the standard normal: centre 0, spread 1, each number independent ($I$, the identity matrix, says "no links between numbers") | |
| $\odot$ | multiply matching entries (element by element) | |
| $z$ | the code the decoder receives this step | 1.1 |
In words: "roll a standard bell-curve number, stretch it by the region's spread, and shift it to the region's centre."
With the numbers: $\sigma = e^{\frac{1}{2} \times (-1.386)} = e^{-0.693} = 0.5$, so $z = 0.5 + 0.5 \times 1.2 = 1.1$. Now the gradients have a route: nudge μ by a little and z moves by the same amount (∂z/∂μ = 1); nudge σ by a little and z moves by ε times as much (∂z/∂σ = ε = 1.2). The symbol ∂ reads "how much this changes when that is nudged".
Level 3: in Python
In Python:
import math
mu, log_var, eps = 0.5, math.log(0.25), 1.2
round(log_var, 3) # → -1.386
# σ = e^(½ log σ²)
sigma = math.exp(0.5 * log_var)
round(sigma, 3) # → 0.5
# z = μ + σ ⊙ ε
z = mu + sigma * eps
round(z, 3) # → 1.1
# nudge μ, then σ, by a tiny h and watch z: ∂z/∂μ = 1 and ∂z/∂σ = ε
h = 1e-6
round((mu + h + sigma * eps - z) / h, 3) # → 1.0
round((mu + (sigma + h) * eps - z) / h, 3) # → 1.2
flowchart LR subgraph A["Sampling directly: the gradient stops"] direction LR m1["μ, σ"] --> s1["draw z from N(μ, σ²)<br/>a dice roll"] --> d1["decoder"] --> l1["loss"] l1 -. "no route back<br/>through a dice roll" .-> s1 end subgraph B["Reparameterized: the gradient flows"] direction LR e2["ε from N(0, I)<br/>just another input"] --> s2["z = μ + σ·ε<br/>plain arithmetic"] m2["μ, σ"] --> s2 --> d2["decoder"] --> l2["loss"] l2 -. "∂z/∂μ = 1, ∂z/∂σ = ε" .-> m2 end
Reading it: both rows produce the same kind of random code: a draw from a bell curve centred on μ with spread σ. In the top row the randomness sits between the encoder's outputs and the loss, and the dotted arrow of backpropagation has nowhere to go. In the bottom row the randomness comes in from the side as ε, an input like a pixel, and everything from μ and σ to the loss is ordinary arithmetic the chain rule can pass through. The encoder learns because of this one rearrangement.
Why it matters in practice. The same trick lets gradients pass through
random choices elsewhere too, for example in the noisy steps of diffusion
models (primer.ml.generative.diffusion) and in reinforcement learning with
continuous actions. It is a large part of why the VAE paper mattered.
In code: reparameterize is the formula; in VAE.loss_and_gradients the lines under "Through z = μ + σ·ε" are the two gradients above, and a test checks them against finite differences.
The KL penalty: rent for every region
Everyday picture. The rent from the everyday picture has a formal name:
the KL divergence (Kullback-Leibler divergence) from the region to the
standard bell curve, a measure of how different two distributions are that
is 0 only when they match (primer.ml.training_stages uses it for
distillation). For a bell-curve region and the standard bell curve, it has a
short closed form:
Level 3: the formula and its symbols
$$ \mathrm{KL}\bigl(\mathcal{N}(\mu, \sigma^2) \,\big\Vert\, \mathcal{N}(0, 1)\bigr) = \frac{1}{2} \sum_{j=1}^{k} \bigl( \mu_j^2 + \sigma_j^2 - \log \sigma_j^2 - 1 \bigr) $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $\mathrm{KL}(q \,\Vert\, p)$ | how much distribution $q$ differs from distribution $p$; 0 when they match | |
| $\mathcal{N}(\mu, \sigma^2)$ | the picture's region: a bell curve with centre μ and variance σ² | $\mathcal{N}(0.5, 0.25)$ |
| $\mathcal{N}(0, 1)$ | the standard bell curve every region is pulled towards | |
| $k$ | how many numbers in the code | 1 (the lesson's network uses 2) |
| $j$ | a counter over the code numbers | |
| $\mu_j^2$ | rent for sitting away from the centre: 0 at μ = 0, growing on both sides | 0.25 |
| $\sigma_j^2 - \log \sigma_j^2$ | rent on size: smallest (exactly 1) when σ = 1, huge as σ shrinks to a pin, growing if it bloats | $0.25 + 1.386$ |
| $-1$ | shifts the total so a perfect match costs exactly 0 | |
| $\frac{1}{2}$ | a constant that falls out of the bell curve's formula |
In words: "for each code number, add the squared distance of its centre from 0, plus its variance, minus the log of its variance, minus 1; add those up over the code numbers and halve."
With the numbers: for μ = 0.5 and σ = 0.5, ½ (0.25 + 0.25 − (−1.386) − 1) = ½ × 0.886 = 0.443. It splits into two rents: the centre alone costs ½ × 0.25 = 0.125, the shrunken size alone ½ (0.25 + 1.386 − 1) = 0.318, and 0.125 + 0.318 = 0.443. A second code number that already has μ = 0 and σ = 1 adds ½ (0 + 1 − 0 − 1) = 0, so the 2-number total is still 0.443. Shrink σ to a pin of 0.05 and the rent jumps to 2.62.
Level 3: in Python
In Python:
import math
def kl(mu, var):
return 0.5 * (mu ** 2 + var - math.log(var) - 1)
# one code number with μ = 0.5, σ = 0.5 (σ² = 0.25)
round(kl(0.5, 0.25), 3) # → 0.443
# the rent for the centre alone, then for the size alone
round(0.5 * 0.5 ** 2, 3) # → 0.125
round(kl(0.0, 0.25), 3) # → 0.318
# a code number that already is the standard normal pays nothing
kl(0.0, 1.0) # → 0.0
# Σ over j: two code numbers, their rents add
round(kl(0.5, 0.25) + kl(0.0, 1.0), 3) # → 0.443
# shrink σ to a pin (σ = 0.05) and the rent jumps
round(kl(0.5, 0.05 ** 2), 2) # → 2.62
Reading it: the left curve holds the spread at σ = 1 and moves the centre: a bowl with its bottom at μ = 0, so sitting at 0.5 costs 0.125 and sitting at 3 costs 4.5. The right curve holds the centre at 0 and changes the spread. It is 0 at σ = 1 (the dashed line), rises gently if the region bloats, and shoots up as σ heads towards 0. That steep wall on the left is what stops the encoder from shrinking every region to a pin and turning back into a plain autoencoder with holes. The two red dots add up to the worked example's 0.443.
In code: kl_to_standard_normal is the formula, summed over the code numbers; a test checks it against a brute-force average over 400,000 samples.
The whole VAE loss
Level 3: the formula and its symbols
$$ L = \bigl\lVert x - \hat{x} \bigr\rVert^2 + \beta \cdot \mathrm{KL}\bigl(q(z \mid x) \,\big\Vert\, \mathcal{N}(0, I)\bigr) $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $x$ | the picture | (2, 3) from the line example |
| $\hat{x}$ | the decoder's rebuild from the sampled code $z$ | (1.6, 3.2) |
| $\lVert x - \hat{x} \rVert^2$ | the rebuild error, as in the plain autoencoder | 0.2 |
| $q(z \mid x)$ | the region the encoder draws for picture $x$: read "q of z given x" | $\mathcal{N}(0.5, 0.25)$ |
| $\beta$ | "beta": how much one unit of KL rent costs, in units of rebuild error | 0.3 |
| KL(…) | the rent from the previous section | 0.443 |
In words: "rebuild well, but pay β for every unit by which your regions differ from the standard bell curve."
With the numbers: 0.2 + 0.3 × 0.443 = 0.2 + 0.133 = 0.333.
Level 3: in Python
In Python:
rebuild, kl_value, beta = 0.2, 0.443, 0.3
# L = ‖x − x̂‖² + β · KL
round(rebuild + beta * kl_value, 3) # → 0.333
Where does this come from? The VAE paper derives the loss (with β = 1) as the negative of the evidence lower bound (ELBO): a quantity that is guaranteed never to exceed the log-probability the model gives to the real pictures, so pushing the bound up pushes that probability up. Choosing squared error amounts to assuming each pixel is the decoder's output plus bell-curve noise of spread s, and working that through gives β = 2s². This lesson uses β = 0.3, which assumes pixel noise of spread about 0.39.
Reading it: the same 200 pictures as the plain autoencoder's map, now placed at their centres μ, each with its fuzzy region shaded around it (regions are drawn one σ wide). The dashed circles are radius 1 and 2 of the standard bell curve. Compare the axes with the plain autoencoder's: the codes now sit around 0 with a spread of about 1.1 in each direction, instead of wandering from −24 to 14. The strands are still there (the code still knows direction and offset), but they are packed side by side and the regions (σ about 0.14) overlap along each strand, so there is far less empty field left.
In code: VAE.loss_and_gradients computes both terms and their gradients; trained_vae trains the lesson's VAE at a chosen β.
Sampling new strokes
Everyday picture. Once the map is packed, throw a dart at the middle of it and you land in a neighbourhood, not a field. Generating is exactly that: throw away the encoder, draw a code from the standard bell curve, and decode.
Tiny example. Draw the same code as before, z = (1.2, −0.9), and decode it with the VAE. Out comes a horizontal stroke through the centre of the picture, with squared distance about 0.06 from a real stroke in the training set. The plain autoencoder turned that same code into a smear 7.5 away. Over 500 draws the VAE's median distance is about 0.56 against the plain autoencoder's 3.
flowchart LR N["draw z from N(0, I)<br/>2 random numbers"] --> D["decoder g"] --> X["a new picture<br/>never seen in training"] ENC["encoder"] -. "not used when generating" .-> N
Reading it: generating uses only the right half of the network. The encoder's job was done during training: it taught the decoder, through the KL rent, that codes live where the standard normal puts them. That is why the first box can draw from $\mathcal{N}(0, I)$ with confidence. The dotted arrow is a reminder that nothing flows from the encoder here.
Reading it: both panels decode the same 24 random codes. The plain autoencoder (left) gives blobs, broken lines and a few strokes by luck. The VAE (right) gives strokes in all four directions at many offsets, plus a few soft blends of two directions. Those blends are the VAE's honest weak spot: with only 2 numbers to hold four directions, some codes sit on the border between two strands, and the decoder hedges between them. About a third of the VAE's samples still cross the junk line, most of them blends like these.
Reading it: every tile is the VAE's decoding of one point on an even grid over the middle of the code map. Read along any row or column: the stroke changes a little at a time, sliding or turning, with no sudden jumps to junk. Whole regions of the map belong to one direction, and the borders between them are where the blends live. This is what a latent space means: a space of codes where position means something. That smoothness is what lets you edit or explore a picture by moving its code.
Interpolation: walking between two codes
Everyday picture. A morph between two faces in a film: every in-between frame should look like a face, not a double exposure.
Tiny example. Encode two pictures to codes $z_a$ and $z_b$, step along the straight line between them, and decode every stop:
Level 3: the formula and its symbols
$$ z_t = (1 - t)\, z_a + t\, z_b, \qquad 0 \le t \le 1 $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $z_a, z_b$ | the codes of the two pictures at the ends of the walk | $(-1, 0.5)$ and $(1, 1.5)$ |
| $t$ | how far along the walk: 0 at $z_a$, 1 at $z_b$ | 0.25 |
| $z_t$ | the code at that point of the walk | $(-0.5, 0.75)$ |
In words: "take a share $1 - t$ of the first code and a share $t$ of the second, and add them."
With the numbers: at $t = 0.25$, $z_t = 0.75 \times (-1, 0.5) + 0.25 \times (1, 1.5) = (-0.75 + 0.25,\ 0.375 + 0.375) = (-0.5, 0.75)$. The walk in the figure uses nine stops, $t = 0, 1/8, 2/8, \ldots, 1$.
Level 3: in Python
In Python:
z_a, z_b = [-1.0, 0.5], [1.0, 1.5]
# z_t = (1 − t) z_a + t z_b, a quarter of the way along
t = 0.25
[(1 - t) * a + t * b for a, b in zip(z_a, z_b)] # → [-0.5, 0.75]
# the nine stops of a walk
[i / 8 for i in range(9)] # → [0.0, 0.125, 0.25, 0.375, 0.5, 0.625, 0.75, 0.875, 1.0]
Reading it: both rows walk between the same two vertical strokes, one left of centre and one right of it. The plain autoencoder's walk (top) leaves its strand and crosses a hole: the middle frames are a smeared cross, 3.88 away from any real stroke. The VAE's walk (bottom) takes a surprising route, through diagonal strokes, because in its 2-number map the diagonal strand sits between those two points. But every stop is a believable stroke, never more than 0.70 from a real one. Averaged over 60 random pairs, the VAE's biggest jump from one frame to the next is also smaller (about 1.6 against 2.4).
Why it matters in practice. Smooth latent spaces are what made "move this slider to add a smile" demos possible, and interpolation is still a standard sanity check for any learned code: if the walk is full of junk, the space has holes.
In code: generate draws and decodes; interpolate encodes two pictures and decodes the walk between them; largest_step measures its biggest jump.
The trade-off: β, blur and collapse
Everyday picture. Go back to the rent. Set it too low and every region shrinks to a pin far from the centre: rebuilds are sharp, but the holes come back. Set it too high and every region moves to the centre and swells to the full bell curve. Then every picture's region is the same region, the code carries no information at all, and the decoder can do nothing better than draw the average of all the strokes: a grey smudge. That failure is called posterior collapse.
Tiny example. The lesson trains six VAEs, with β from 0.01 to 3:
| β | rebuild error | KL (information in the code) | median sample distance to a real stroke |
|---|---|---|---|
| 0.01 | 0.16 | 10.3 | 2.40 |
| 0.1 | 0.23 | 5.7 | 0.87 |
| 0.3 | 0.31 | 4.4 | 0.56 |
| 1 | 0.92 | 2.9 | 0.81 |
| 3 | 6.08 | 0.0 | 4.73 |
(beta_sweep produces this table, including β = 0.03, and demo prints it.)
The samples are best in the middle. At β = 3 the KL is zero: the regions are the standard bell curve itself, and the rebuild error of 6.08 is what you get by drawing the average stroke every time.
Why VAE pictures are blurry
Even at a good β the regions overlap, so one code can stand for several different pictures. What should the decoder draw for it? Take a single pixel that is black (1) in one of those pictures and white (0) in the other, equally likely. Squared error has a clear answer:
Level 3: the formula and its symbols
$$ \underset{c}{\arg\min}\ \mathbb{E}\bigl[(x - c)^2\bigr] = \mathbb{E}[x] $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $x$ | the pixel's true value, which could be either picture's | 0 or 1, equally likely |
| $c$ | the decoder's guess for that pixel | 0, 0.5 or 1 |
| $(x - c)^2$ | the squared error of the guess | |
| $\mathbb{E}[\ldots]$ | expected value: the average over the possibilities, each weighted by its chance | |
| $\arg\min_c$ | "the $c$ that makes what follows as small as possible" |
In words: "the guess with the smallest average squared error is the average of the possibilities."
With the numbers: guess black (1): error 1 half the time, 0 otherwise, average 0.5. Guess white (0): also 0.5. Guess grey (0.5): error 0.25 either way, average 0.25. Grey wins, and grey is the average of 0 and 1.
Level 3: in Python
In Python:
outcomes = [0.0, 1.0]
def expected_error(c):
return sum((x - c) ** 2 for x in outcomes) / len(outcomes)
# E[(x − c)²] for three guesses: white, grey, black
[expected_error(c) for c in (0.0, 0.5, 1.0)] # → [0.5, 0.25, 0.5]
# E[x]: the average outcome, which is the winning guess
sum(outcomes) / len(outcomes) # → 0.5
Across a whole picture, "average the pictures this code might mean" is a blur. The more the regions overlap, the more pictures each code must stand for, and the blurrier the output.
Reading it: on the left, β grows along a logarithmic axis. The blue rebuild error stays low, then climbs steeply past β = 1. The green KL, the information the code carries, falls steadily and hits zero at β = 3. The red line, how far random samples are from real strokes, is a U: high at small β (holes), lowest near β = 0.3, high again at large β (blur and collapse). On the right are eight samples at each β. The top rows have sharp ink in junk shapes. The middle rows are clean strokes. The β = 1 row is softer. The bottom row is eight copies of the same grey smudge: the collapse.
Why it matters in practice. β is a real dial, and the name β-VAE
(Higgins and colleagues, 2017) comes from turning it up to get codes whose
numbers line up with separate factors, like direction and offset here. The
blur is the VAE's best-known weakness. Modern systems fix it by adding a
second judge of realism to the loss (an adversarial loss, the trick behind
primer.ml.generative.gans, as in VQGAN), or by letting a stronger generator
do the generating while the autoencoder only compresses.
In code: beta_sweep trains and measures one VAE per β; expected_squared_error is the grey-pixel calculation.
Where autoencoders live today
Everyday picture. The autoencoder became the zip format for images and sound that other generative models work inside.
Tiny example. Stable Diffusion's autoencoder turns a 512 × 512 colour image (512 × 512 × 3 = 786,432 numbers) into a 64 × 64 × 4 grid of codes (16,384 numbers): 48 times fewer. The diffusion model does all its work on that small grid, and the decoder paints full-size pixels only at the very end. Its KL weight is tiny, so it is mostly a compressor, with just enough rent to keep the codes well-behaved.
A VQ-VAE (vector-quantized VAE) goes one step further: it snaps each code vector to the nearest entry of a learned codebook, so an image becomes a grid of whole numbers, the entries' positions in the codebook. Those numbers are tokens, and a transformer can read and write them exactly as it reads and writes words. Neural audio codecs do the same for sound.
flowchart LR subgraph LD["Latent diffusion"] direction LR I1["image<br/>512 × 512 × 3"] --> E1["VAE encoder"] --> L1["latent<br/>64 × 64 × 4"] L1 --> DF["diffusion model<br/>works here"] --> D1["VAE decoder"] --> O1["image"] end subgraph TK["Tokenizer for a transformer"] direction LR I2["image or audio"] --> E2["encoder"] --> Q["snap each vector to<br/>nearest codebook entry"] Q --> T["grid of token ids"] --> TR["transformer"] TR --> D2["decoder"] --> O2["image or audio"] end
Reading it: in the top row the autoencoder is the outer shell: its
encoder shrinks the image 48-fold on the way in, its decoder restores it on
the way out, and the expensive generator in the middle never touches a
pixel. Training and sampling run far faster on 16,384 numbers than on
786,432, which is what made high-resolution diffusion affordable (see
primer.ml.generative.diffusion). In the bottom row the snapping step turns
continuous codes into token ids, which is how images and audio can enter
and leave a language-model-style transformer (see
primer.ml.generative.multimodal).
Why it matters in practice. When a model "generates an image", there is very often an autoencoder at both ends of the pipeline. Its quality sets a ceiling: whatever detail the decoder cannot rebuild, no generator working in its latent space can produce.
In 20 seconds
- Autoencoder: an encoder squeezes data through a narrow code (the bottleneck) and a decoder rebuilds it; training minimises the rebuild error. With no bends it learns exactly what PCA learns; with bends it can follow curved data far better.
- No generation from a plain autoencoder: its codes land wherever training put them, with holes in between, so random codes decode to junk.
- VAE: the encoder outputs a region (μ, σ); the reparameterization trick z = μ + σ·ε samples from it in a way gradients can pass through; a KL penalty pulls every region towards the standard normal, so drawing z ~ N(0, I) and decoding generates new data.
- The trade-off: β weighs the KL. Too small brings back holes, too large blurs and finally collapses; squared error averages over overlapping possibilities, which is why VAE output is blurry.
- Today: VAEs compress images for latent diffusion, and VQ-VAEs turn images and audio into tokens for transformers.
Self-test questions
What is the bottleneck for? What goes wrong if the code is as wide as the input? It forces the network to keep only what matters: with 2 numbers for 64 pixels, it must find the few facts that actually vary. If the code is as wide as the input, the network can learn to copy the pixels straight through, rebuild perfectly and learn nothing useful, unless something else (noise on the input, a penalty on the code) stops it.
Why can't you generate new pictures by decoding random codes from a plain autoencoder? Its loss only ever sees the codes of real pictures, so it says nothing about where codes should live or what lies between them. The codes end up in an arbitrary range with holes between clusters, and a random code usually lands in a hole, where the decoder's output is junk.
What problem does the reparameterization trick solve? Backpropagation needs every step between the weights and the loss to be something you can differentiate, and a random draw is not. Writing the draw as z = μ + σ·ε with the noise ε as a separate input turns sampling into arithmetic, so gradients reach μ and σ, and through them the encoder.
What do the two terms of the VAE loss each want, and why do you need both? The rebuild term wants regions small and far apart so every picture is decoded precisely. The KL term wants every region to be the standard bell curve. Without KL you get a plain autoencoder with holes; without the rebuild term every region collapses onto the bell curve and the code carries nothing. The balance gives a packed, smooth code space that still tells pictures apart.
Work out the KL penalty for one code number with μ = 0 and σ = 2. ½ (0 + 4 − log 4 − 1) = ½ (3 − 1.386) = 0.807. A region that is too big pays rent too, though less steeply than one that is too small.
Why are VAE samples blurry? Regions overlap, so one code can stand for several pictures, and under squared error the best single answer is their average. Averaged pictures are blurred pictures. Raising β increases the overlap and the blur.
What is posterior collapse? When the KL rent outweighs what the code saves in rebuild error, the encoder makes every region the standard bell curve, the code carries no information, and the decoder outputs the same average picture for every code. In this lesson that happens at β = 3.
Why does latent diffusion run inside an autoencoder's code space instead of on pixels? The code is many times smaller (48 times for Stable Diffusion's 512 × 512 images) and keeps what matters to the eye, so the expensive, many-step generator is far cheaper to train and run. The decoder turns the result back into full-size pixels once, at the end.
The papers behind this lesson
- Hinton & Salakhutdinov, Reducing the Dimensionality of Data with Neural Networks (Science, 2006): https://doi.org/10.1126/science.1127647. Showed that deep autoencoders, once they could be trained, compress data far better than PCA.
- Kingma & Welling, Auto-Encoding Variational Bayes (2013): https://arxiv.org/abs/1312.6114. Introduced the variational autoencoder, the reparameterization trick and the ELBO loss with its closed-form Gaussian KL. Annotated companion
- Rezende, Mohamed & Wierstra, Stochastic Backpropagation and Approximate Inference in Deep Generative Models (2014): https://arxiv.org/abs/1401.4082. Developed the same idea independently at the same time, showing how to backpropagate through random sampling.
- Burgess et al., Understanding disentangling in β-VAE (2018): https://arxiv.org/abs/1804.03599. Explains why turning up β pushes the code's numbers to line up with separate factors of the data, and what it costs in rebuild quality.
- van den Oord, Vinyals & Kavukcuoglu, Neural Discrete Representation Learning (2017): https://arxiv.org/abs/1711.00937. Introduced the VQ-VAE, which snaps codes to a learned codebook and so turns images and audio into tokens. Annotated companion
- Rombach et al., High-Resolution Image Synthesis with Latent Diffusion Models (2021): https://arxiv.org/abs/2112.10752. Ran diffusion inside a lightly regularized autoencoder's code space, the design behind Stable Diffusion.
Further reading
- Kingma & Welling, An Introduction to Variational Autoencoders (2019): https://arxiv.org/abs/1906.02691
- Carl Doersch, Tutorial on Variational Autoencoders (2016): https://arxiv.org/abs/1606.05908
- Goodfellow, Bengio & Courville, Deep Learning, chapter 14, Autoencoders: https://www.deeplearningbook.org/contents/autoencoders.html
- Esser, Rombach & Ommer, Taming Transformers for High-Resolution Image Synthesis (VQGAN, 2020): https://arxiv.org/abs/2012.09841
- PyTorch's VAE example: https://github.com/pytorch/examples/tree/main/vae
1r""" 2# Autoencoders and VAEs 3 4Run: `python -m primer.ml.generative.autoencoders` 5 6This lesson builds on the two-layer network and training loop of 7`primer.ml.neural_net` and the Adam optimizer of `primer.ml.optimizers`; 8`primer.notation` explains every symbol from zero. 9 10## Level 1: The practitioner's guide 11 12**In one sentence.** An autoencoder is a pair of networks trained to squeeze 13data through a narrow code and rebuild it, which makes it a learned lossy 14compressor; a variational autoencoder (VAE) also shapes the code space so 15that a random code decodes to something sensible, which makes it a 16generator and, far more often today, the compressed space that bigger 17generators (diffusion models, transformers) work inside. 18 19**When you need it.** Three tells. You have unlabelled data and want a 20compact, meaningful representation of each item (a fingerprint for search, 21a small input for another model, a way to flag the items that don't fit): 22that is a plain autoencoder. You want to generate or edit new items and 23need a smooth space where nearby codes mean similar outputs: that is a VAE. 24You are building or running an image, audio or video generator: there is 25almost certainly an autoencoder at both ends of it, and its settings 26(downsampling factor, latent channels, scale factor, precision) are yours 27to get right. You don't need one for compression that must be exact (a zip 28file is lossless; an autoencoder never is), and you rarely train one for 29images or audio yourself any more: pretrained ones are downloadable, and 30their quality took a great deal of data to reach. The number that shows 31the naive path failing: in this lesson a plain autoencoder rebuilds 8 × 8 32pen strokes from two numbers with an error of 0.07 (99% of the picture kept, 33against 50% for PCA), yet decoding random codes drawn from the standard bell 34curve gives junk 76% of the time. Compression is not generation. 35 36**Your options.** From the cheapest to the most capable: 37 38| Option | What it does | What it guarantees | What it costs | Where it lives | 39|---|---|---|---|---| 40| PCA | Fits a flat sheet through the data; an item's code is where it lands on the sheet | The best any flat code can do under squared error; no training loop | One matrix decomposition; poor rebuilds of curved data (50% kept here) | A library call | 41| Plain autoencoder | A bent encoder and decoder trained only to rebuild the input | Far better rebuilds on curved data (99% kept here); a code that flags anomalies and can denoise | A training run; a code space with holes, so no generation | Your training loop | 42| VAE | The encoder emits a fuzzy region per item and pays rent for straying from the standard bell curve | Random codes decode sensibly (36% junk here, against 76%); smooth interpolation | Blurrier rebuilds; a β to tune; posterior collapse if you overdo it | Your training loop, or a pretrained one | 43| Latent autoencoder for diffusion | A VAE with a tiny KL weight; the diffusion model generates inside its code space | 48 times fewer numbers for a 512 × 512 image, so training and sampling become affordable | A ceiling on detail set by the decoder; scaling, precision and memory settings to respect | Downloaded with the diffusion model | 44| VQ-VAE and VQGAN | Snaps each code vector to the nearest entry of a learned codebook, so an item becomes a grid of token ids | Tokens a transformer reads and writes like words; sharper output when an adversarial loss is added (VQGAN) | A codebook to keep in use; a discrete space with no straight-line interpolation | A pretrained tokenizer | 45| Neural audio codec | The same recipe for sound: encoder, residual quantizer, decoder, reconstruction plus adversarial losses | Speech and music at 3 to 18 kbit/s, streamable in real time | A model at both ends of the wire | SoundStream, EnCodec | 46 47**How to choose.** Start from what you want the code for. 48 49- Compact features for search, clustering or a downstream model, and no 50 generation: try PCA first (Level 2 shows it is exactly the autoencoder 51 with no bends). Train an autoencoder when the data is curved and PCA's 52 rebuilds are poor, as they are here. 53- Flagging oddities (fraud, a failing machine): a plain autoencoder trained 54 on normal data, with an alarm on rebuild error. 55- Generating or editing new items in a small domain, or sliders that mean 56 something: a VAE, with β chosen by looking at samples, not at the loss. 57- Generating images, audio or video at real resolution: don't generate 58 with the VAE. Use it as the compressor and let diffusion or a transformer 59 do the generating, and download the autoencoder that generator was 60 trained with, because the pair is matched. 61- Feeding images or audio into a language-model-style transformer: a VQ 62 tokenizer or a neural codec. 63- Whatever you pick, look at the rebuilds before anything else. The 64 decoder's rebuild quality is the ceiling of every generator that works in 65 its space; no generator can produce detail the decoder cannot paint. 66 67**What it costs.** Training cost is dominated by data: the autoencoder must 68see enough of the domain to rebuild it. At run time it is cheap: one 69encoder pass on the way in and one decoder pass on the way out, against the 70many passes of the generator between them. That is the economics of latent 71diffusion. Rombach et al. tried downsampling factors 72from 1 (raw pixels) to 32 and found factors 4 and 8 the sweet spot; after 73two million training steps the pixel-space model trailed the factor-8 model 74by 38 FID points, and on inpainting the latent models ran at least 2.7 times 75faster. Memory: at high resolution the decoder's activations fill the GPU, 76which is why the diffusers `AutoencoderKL` offers tiled encoding and 77decoding (constant memory, at the risk of faint tile seams) and runs the 78SDXL autoencoder in float32 by default. Quality has a measured ceiling: the 79Stable Diffusion 3 paper reports that widening the latent from 4 to 8 to 16 80channels drops reconstruction FID from 2.41 to 1.56 to 1.06 and raises PSNR 81from 25.12 to 26.40 to 28.62, which is why newer models carry 16-channel 82latents at the price of a harder generation task. β is a dial with a bottom: 83in this lesson, samples land nearest real strokes at β = 0.3 (median 84distance 0.56) and further at both β = 0.01 (2.40) and β = 3 (4.73). Tokens 85cost context: DALL-E's discrete VAE turns a 256 × 256 image into 32 × 32 = 861,024 tokens from a codebook of 8,192, cutting the transformer's context 87192-fold. 88 89**What breaks.** 90 91- **Holes.** A plain autoencoder's codes land wherever training put them 92 (from −24 to 14 here), with empty fields between; 44% of random codes 93 drawn even from inside that range decode to junk. If you need to sample, 94 you need the KL rent: use a VAE. 95- **Blur.** Under squared error the best guess for an uncertain pixel is 96 the average, so a VAE whose regions overlap paints averages. Lower β for 97 sharper rebuilds, or add an adversarial loss as VQGAN does, or stop asking 98 the VAE to generate and let a stronger model do it in its space. 99- **Posterior collapse.** Raise β too far and every region becomes the 100 standard bell curve, the code carries nothing, and the decoder emits the 101 same average picture for every code (β = 3 here: KL exactly 0, rebuild 102 error 6.08, one grey smudge). Watch the KL term; zero is a symptom. 103- **Forgetting the latent scale.** Diffusion libraries multiply latents by 104 a scaling factor (0.18215 for Stable Diffusion's autoencoder) so they have 105 unit variance for the generator, and divide it back out before decoding. 106 Skip either step and the generator sees data it never trained on. 107- **Precision.** The Stable Diffusion autoencoders overflow in float16 at 108 high resolution; run them in float32 or use a checkpoint fine-tuned for 109 half precision. 110- **Dead codebook entries.** In a VQ model, entries nothing maps to waste 111 the vocabulary. DALL-E raised its KL weight to 6.6 to promote codebook 112 usage; if your tokens cluster on a few ids, that is the dial. 113- **A mismatched pair.** Latents from one autoencoder decoded by another 114 are junk. Keep the encoder, the generator and the decoder that were 115 trained together. 116 117**In the wild.** Stable Diffusion's KL-regularised autoencoder (8 times 118downsampling, 4 latent channels) ships with every Stable Diffusion model 119and is `AutoencoderKL` in Hugging Face diffusers; Stable Diffusion 3 moved 120to 16 channels. VQGAN (Esser, Rombach and Ommer) adds an adversarial loss 121to a VQ-VAE and puts a transformer over its tokens; the original VQ-VAE 122compressed 128 × 128 images to a 32 × 32 grid over a codebook of 512, about 12342.6 times fewer bits, and generated with a PixelCNN over the grid; DALL-E's 124discrete VAE with 8,192 codes fed its text-to-image transformer. For sound, 125SoundStream (a convolutional encoder and decoder around a residual vector 126quantizer, 3 to 18 kbit/s, and at 3 kbit/s preferred over Opus at 12) and 127EnCodec (a streaming encoder-decoder with a multiscale spectrogram 128adversary, at 24 kHz mono and 48 kHz stereo) are the same recipe at audio's 129scale. Every paper is linked at the end of the lesson. 130 131**Go deeper.** Level 2 builds both halves by hand on 8 × 8 pen strokes, 132shows where a plain autoencoder's holes come from, then adds the VAE's two 133pieces (the reparameterization trick and the KL rent) with numbers you can 134check, and sweeps β to watch holes give way to blur and then collapse. If 135you only needed to choose, you are done. 136 137## Level 2: How it works, from scratch 138 139You phone a friend and describe a picture so they can draw it, but you are 140allowed to say only two numbers. You would agree on a system first: the first 141number says which way the line runs, the second says how far it sits from the 142middle. You squeeze the picture into two numbers, send them, and your friend 143rebuilds it. 144 145An **autoencoder** is a neural network that invents that system by itself. It 146has two halves. The **encoder** squeezes a picture into a few numbers, called 147the **code**. The **decoder** rebuilds the picture from the code. Nobody tells 148it what the numbers should mean. The only instruction is "the rebuild must 149match the original", and to obey it through such a narrow gap the network has 150to discover what really varies in the data. 151 152Then comes a second question. If your friend can draw from *any* two numbers, 153can you invent a new picture by making up two numbers? With a plain 154autoencoder, usually not: most made-up codes were never used for anything, 155and the drawing comes out as a smudge. A **variational autoencoder (VAE)** is 156trained so that a made-up code drawn from a known range draws something 157sensible. That step, from compressing to generating, is what this lesson is 158about, and it opens this part of the primer: GANs 159(`primer.ml.generative.gans`) and diffusion (`primer.ml.generative.diffusion`) 160are other answers to the same question. 161 162## The toy data: pictures of pen strokes 163 164**Tiny example.** Every picture in this lesson is 8 × 8 = 64 pixels holding 165one soft pen stroke. The stroke runs in one of four directions (horizontal, 166vertical, diagonal down, diagonal up) and is shifted up to 2.5 pixels from 167the centre. The ink fades like a bell curve either side of the line: a pixel 168on the line has brightness 1.0, a pixel one pixel away about 0.25, a pixel 169two away almost nothing. 170 171So each picture is 64 numbers, but only **two facts** change from picture to 172picture: the direction (one of four) and the offset (any amount). A good 1732-number code has to rediscover both from the pixels alone. The 200 pictures 174are the top row of the first figure below. 175 176**In code:** `make_strokes` draws the pictures and returns their hidden directions and offsets; `STROKE_KINDS` names the four directions. 177 178## Squeeze and rebuild: the autoencoder 179 180**Everyday picture.** A zip file squeezes a document and gets it back 181exactly. An autoencoder is a *lossy*, *learned* zip: it keeps what matters 182most for this kind of data and lets the rest go. 183 184**Tiny example: a one-number code you can check by hand.** Take points that 185lie on the line y = 2x, such as (1, 2) and (2, 4). Each takes two numbers to 186write down, but one number is enough: how far along the line the point is. 187 188- **Encoder:** code = u · x, where u = (1, 2)/√5 = (0.447, 0.894) is the 189 line's direction scaled to length 1, and "·" is the dot product (multiply 190 matching entries and add; see `primer.notation`). 191- **Decoder:** rebuild = code × u. 192 193For (1, 2): code = 0.447 + 1.789 = 2.236 (which is √5, its distance from the 194origin), and rebuild = 2.236 × (0.447, 0.894) = (1, 2). Nothing lost. 195 196For (2, 3), which is *off* the line: code = (2 + 6)/√5 = 3.578, and 197rebuild = 3.578 × (0.447, 0.894) = (1.6, 3.2), the closest point on the line. 198The part that pointed away from the line, (0.4, −0.2), is gone. That loss is 199exactly what training measures and shrinks: 200 201$$ 202L_{\text{rec}} = \frac{1}{n} \sum_{i=1}^{n} \bigl\lVert x_i - g\bigl(f(x_i)\bigr) \bigr\rVert^2 203$$ 204 205**Symbols** 206 207| Symbol | Meaning here | In the example | 208|---|---|---| 209| $x_i$ | the $i$-th example: a point, or a picture's 64 pixels | $x = (2, 3)$ | 210| $f$ | the encoder: squeezes an example into a code | $f(x) = u \cdot x$ | 211| $f(x_i)$ | the code for example $i$, often written $z_i$ | 3.578 | 212| $g$ | the decoder: rebuilds an example from a code | $g(z) = z\,u$ | 213| $g(f(x_i))$ | the rebuild, often written $\hat{x}_i$ ("x hat") | (1.6, 3.2) | 214| $x_i - \hat{x}_i$ | what the rebuild got wrong, entry by entry | (0.4, −0.2) | 215| $\lVert v \rVert^2$ | squared length of $v$: square every entry and add | $0.4^2 + 0.2^2 = 0.2$ | 216| $\sum_{i=1}^{n}$ | add up over every example | | 217| $n$ | how many examples | 2: (1, 2) and (2, 3) | 218| $L_{\text{rec}}$ | the reconstruction loss: average squared error of the rebuilds | 0.1 | 219 220**In words:** "squeeze each example, rebuild it, measure the squared distance 221between the rebuild and the original, and average over all the examples." 222 223**With the numbers:** (1, 2) rebuilds perfectly, error 0. (2, 3) rebuilds as 224(1.6, 3.2), error 0.4² + (−0.2)² = 0.16 + 0.04 = 0.2. Averaged over the two 225points, $L_{\text{rec}}$ = (0 + 0.2)/2 = 0.1. 226 227**In Python:** 228 229```python 230import math 231# u: the line's direction, scaled to length 1 232u = [1 / math.sqrt(5), 2 / math.sqrt(5)] 233def f(x): 234 return sum(u_m * x_m for u_m, x_m in zip(u, x)) 235def g(z): 236 return [z * u_m for u_m in u] 237def error(x): 238 return sum((a - b) ** 2 for a, b in zip(x, g(f(x)))) 239# f(x): the encoder squeezes (2, 3) to one number 240round(f([2, 3]), 3) # → 3.578 241# g(f(x)): the decoder rebuilds two numbers from it 242[round(v, 2) for v in g(f([2, 3]))] # → [1.6, 3.2] 243# ‖x − x̂‖² for each point: (1, 2) is on the line, (2, 3) is not 244[round(error(x), 2) for x in ([1, 2], [2, 3])] # → [0.0, 0.2] 245# L_rec: the average over n = 2 points 246round((error([1, 2]) + error([2, 3])) / 2, 2) # → 0.1 247``` 248 249The real network has the same two halves, only bent. The encoder is a 250two-layer network like the one built in `primer.ml.neural_net`: 64 pixels 251into 32 hidden numbers (through tanh, which lets it bend), then out to a code 252of 2 numbers. The decoder mirrors it: 2 numbers into 32 hidden, then out to 25364 pixels squashed between 0 and 1 by a sigmoid so they are valid 254brightnesses. Training is the loop from `primer.ml.neural_net` with the Adam 255optimizer from `primer.ml.optimizers`: rebuild every picture, measure 256$L_{\text{rec}}$, send the gradient back through both halves, adjust, 1,500 257times. 258 259```mermaid 260flowchart LR 261 X["picture x<br/>64 pixels"] --> E["encoder f<br/>64 → 32 → 2"] 262 E --> Z["code z<br/>2 numbers"] 263 Z --> D["decoder g<br/>2 → 32 → 64"] 264 D --> XH["rebuild x̂<br/>64 pixels"] 265 X --> L["loss<br/>squared error between x and x̂"] 266 XH --> L 267 L -. "gradients flow back<br/>through decoder, then encoder" .-> E 268``` 269 270**Reading it:** follow the picture left to right. It is 64 numbers wide at 271both ends and only 2 numbers wide in the middle: that narrow middle is the 272**bottleneck**, and it is the whole point. Without it the network could copy 273the pixels straight through and learn nothing. The loss box compares the two 274ends, and the dotted arrow is backpropagation carrying the blame back through 275the decoder and on into the encoder, so both halves learn together. The 276encoder never sees a target code; it learns whatever code the decoder finds 277most useful. 278 279 280 281**Reading it:** the top row is eight real strokes, two of each direction. 282The middle row is what the autoencoder rebuilds from just 2 numbers per 283picture: nearly identical. The average error is about 0.07, against about 2848.8 units of squared ink in a whole stroke, so it keeps over 99% of the 285picture. The bottom row squeezes the same pictures to 2 numbers with PCA and 286rebuilds them: grey smudges, keeping only about half. Same budget of two 287numbers, very different results. The next section explains why. 288 289**In code:** `worked_example_line` is the hand example above; `Autoencoder` holds both halves and its hand-written backward pass in `Autoencoder.loss_and_gradients`; `train` runs Adam; `trained_autoencoder` is the lesson's trained network; `reconstruction_error` measures the rebuilds. 290 291### PCA is the straight-line special case 292 293**Everyday picture.** PCA (principal component analysis, built in 294`primer.ml.embeddings.clustering`) fits a flat sheet through the data and 295records where each point lands on the sheet. That is an autoencoder with no 296bends: the encoder is one matrix multiply, and so is the decoder. 297 298**Tiny example.** Scatter 100 points near the line y = 2x and train exactly 299that no-bend autoencoder, with a 1-number code, by plain gradient descent. 300The direction it learns is (0.447, 0.894), the line's own direction, and its 301rebuild error matches PCA's with one component. Baldi and Hornik proved in 3021989 that this always happens: a linear autoencoder trained on squared error 303lands on the same flat sheet as PCA. 304 305So why did PCA smudge the strokes? Because a stroke sliding across the image 306does not travel in a straight line through pixel space. Take a horizontal 307stroke at the top and another at the bottom. Their average, the point halfway 308along the straight line between them, is *two faint lines*, not one stroke in 309the middle. The real strokes lie on a curved surface in the 64-dimensional 310space of pictures, and a flat sheet can only cut through it. The 311autoencoder's tanh layers let its "sheet" bend to follow the curve. 312 313**Why it matters in practice.** Before generation, autoencoders earned their 314keep as learned compressors. They compress, they denoise (train on noisy 315inputs, ask for clean outputs), and they detect anomalies: an input the 316network rebuilds badly is unlike anything it trained on, which is a standard 317way to flag fraud or a failing machine. 318 319 320 321**Reading it:** each dot is one picture, placed at its 2-number code and 322coloured by its stroke direction, with bigger dots for larger offsets. The 323codes form **strands**: pictures of one direction line up along a curve (now 324and then broken into pieces), and sliding along a strand slides the stroke. 325The network rediscovered both hidden facts without being told either. Now 326look at the axes: the codes run from about −24 to 14. Nothing asked for that 327range; it is an accident of training. The dashed circle near 0 is where a 328"random" code from the bell curve would usually land, and it catches almost 329none of the strands. Hold on to that for the next section. 330 331**In code:** `train_linear_autoencoder` is the no-bend autoencoder; `pca_reconstruction` rebuilds from the top principal components. 332 333## Why a plain autoencoder cannot generate: holes 334 335**Everyday picture.** Imagine a town where houses were built only along four 336winding roads. Pick a random spot inside the town limits and you will most 337likely land in a field. Ask the decoder to draw the picture that "lives" in 338that field, and it answers anyway, with whatever its weights happen to 339produce, because nothing in training ever asked it about that spot. 340 341**Tiny example.** The obvious way to make up a random code is to draw each 342number from the **standard normal distribution**, the bell curve centred on 3430 with a spread (standard deviation) of 1: most draws fall between −1 and 1, 344almost all between −3 and 3. Written $\mathcal{N}(0, I)$ for several numbers 345at once, it means "draw each number from that bell curve, independently". 346 347Draw the code z = (1.2, −0.9) and decode it with the plain autoencoder. The 348result is a bright smear with about three times the ink of any real stroke; 349its squared distance to the nearest real stroke is about 7.5. Do that 500 350times and the median distance is about 3, while three quarters of the draws 351land further than 1.0 from every real stroke. That 1.0 is this lesson's line 352for **junk**: about a ninth of a whole stroke's squared ink. 353 354To be fair to the autoencoder, draw instead from the box that holds its own 355codes, from −24 to 1.6 across and −17 to 14 up. Still about 44% of those 356codes decode to junk. They fall in the **holes** between the strands. 357 358 359 360**Reading it:** on the left, black dots are the codes of real strokes and the 361dashed rectangle is the box around them. Every blue circle is a random code 362that decoded to something close to a real stroke; every red cross decoded to 363junk. The red crosses sit in the open spaces between strands, the blue 364circles near them. On the right are the decoded pictures themselves: the 12 365worst (top two rows) are black blobs no pen would draw, and the 12 best 366(bottom) are clean strokes, because those random codes happened to land on 367a strand. 368 369**Why it matters in practice.** A compressor is not a generator. To sample 370new data you need a code space with a **known shape** that you can draw from, 371**no holes** inside that shape, and **smoothness**, so nearby codes decode to 372similar pictures. A plain autoencoder promises none of these, because its 373loss only ever looks at codes of real pictures. 374 375**In code:** `generate` decodes codes drawn from the standard normal; `codes_in_box` draws from a model's own code range; `distance_to_data` measures how far each image is from the nearest real stroke, and `JUNK_DISTANCE` is the line between plausible and junk. 376 377## The VAE: encode to a fuzzy region, not a point 378 379**Everyday picture.** Instead of pinning each picture to one exact spot on 380the map, the encoder draws a small fuzzy circle: "somewhere around here". 381During training the decoder is handed a random point from inside that circle, 382so it must draw the right picture from anywhere nearby. A whole neighbourhood 383now decodes sensibly, not just one pin. A second rule stops the encoder from 384cheating: every circle pays **rent**, more the further it sits from the 385middle of the map and more if it shrinks towards a pin. Circles crowd 386towards the centre and overlap, and the fields between the roads fill in. 387 388**Tiny example.** One picture, one code number. The encoder says 389μ = 0.5 and σ = 0.5: "about 0.5, give or take 0.5". This step, a bell-curve 390roll gives ε = 1.2, so the decoder is handed z = 0.5 + 0.5 × 1.2 = 1.1. Next 391step the roll is ε = −0.4 and the decoder gets z = 0.5 − 0.2 = 0.3. Both must 392rebuild the same picture. 393 394```mermaid 395flowchart LR 396 X["picture x"] --> E["encoder"] 397 E --> MU["μ: centre of the region"] 398 E --> LV["log σ²: size of the region"] 399 EPS["ε drawn from N(0, I)"] --> Z["z = μ + σ·ε"] 400 MU --> Z 401 LV --> Z 402 Z --> D["decoder"] --> XH["rebuild x̂"] 403 XH --> R["rebuild error<br/>‖x − x̂‖²"] 404 X --> R 405 MU --> KL["KL rent<br/>pulls regions to N(0, I)"] 406 LV --> KL 407 R --> LOSS["loss = rebuild + β · KL"] 408 KL --> LOSS 409``` 410 411**Reading it:** compare it with the plain autoencoder's diagram. The encoder 412now has two outputs per code number: a centre μ and a size, given as 413log σ² (the logarithm of the variance, used because it can be any number 414while σ itself must stay positive). The box z = μ + σ·ε is where the fuzzy 415region becomes one concrete code: ε is fresh random noise every step. The 416loss has two parts. The rebuild error, as before, wants each region small and 417distinct so the decoder knows exactly which picture it came from. The KL 418rent, fed straight from μ and log σ², wants every region to look like the 419standard bell curve. Training settles on a compromise between them, and β 420sets the exchange rate. 421 422**In code:** `VAE` is the `Autoencoder` with a two-headed encoder; `VAE.encode_distribution` returns μ and log σ², and `VAE.encode` returns just the centre μ, the best single code for a picture. 423 424### The reparameterization trick: moving the dice outside 425 426**Everyday picture.** To pick a random seat in a row, you could close your 427eyes and point. If someone then moves the row one seat to the left, you have 428no idea how your pick would have changed. Or you could roll a die for "how 429many seats from the middle" and count from wherever the middle is. Now if the 430row moves one seat left, your seat moves exactly one seat left. Same kind of 431random seat, but you can say how it responds to the row moving. 432 433Training needs exactly that. Backpropagation asks of every step: "if I nudge 434this number, how does the loss change?" (that rate of change is the 435**gradient**, or **derivative**; see `primer.notation`). A raw dice roll has 436no answer, so the gradient would stop at the sampling step and the encoder 437would never learn. The **reparameterization trick** rolls the dice first, as 438an ordinary input ε, and builds the code with arithmetic: 439 440$$ 441z = \mu + \sigma \odot \varepsilon, \qquad \varepsilon \sim \mathcal{N}(0, I), \qquad \sigma = e^{\frac{1}{2} \log \sigma^2} 442$$ 443 444**Symbols** 445 446| Symbol | Meaning here | In the example | 447|---|---|---| 448| $\mu$ | the centre of the picture's region, from the encoder | 0.5 | 449| $\log \sigma^2$ | the natural logarithm of the region's variance, from the encoder; "log of y" is the power you raise $e$ to in order to get y | $\log 0.25 = -1.386$ | 450| $e$ | Euler's number, ≈ 2.718 | | 451| $\sigma$ | the spread of the region (standard deviation); $e^{\frac{1}{2}\log\sigma^2}$ undoes the log and the square | 0.5 | 452| $\varepsilon$ | "epsilon": random noise, one number per code number | 1.2 | 453| $\sim$ | "is drawn from" | | 454| $\mathcal{N}(0, I)$ | the standard normal: centre 0, spread 1, each number independent ($I$, the identity matrix, says "no links between numbers") | | 455| $\odot$ | multiply matching entries (element by element) | | 456| $z$ | the code the decoder receives this step | 1.1 | 457 458**In words:** "roll a standard bell-curve number, stretch it by the region's 459spread, and shift it to the region's centre." 460 461**With the numbers:** $\sigma = e^{\frac{1}{2} \times (-1.386)} = e^{-0.693} 462= 0.5$, so $z = 0.5 + 0.5 \times 1.2 = 1.1$. Now the gradients have a route: 463nudge μ by a little and z moves by the same amount (∂z/∂μ = 1); nudge σ by a 464little and z moves by ε times as much (∂z/∂σ = ε = 1.2). The symbol ∂ reads 465"how much this changes when that is nudged". 466 467**In Python:** 468 469```python 470import math 471mu, log_var, eps = 0.5, math.log(0.25), 1.2 472round(log_var, 3) # → -1.386 473# σ = e^(½ log σ²) 474sigma = math.exp(0.5 * log_var) 475round(sigma, 3) # → 0.5 476# z = μ + σ ⊙ ε 477z = mu + sigma * eps 478round(z, 3) # → 1.1 479# nudge μ, then σ, by a tiny h and watch z: ∂z/∂μ = 1 and ∂z/∂σ = ε 480h = 1e-6 481round((mu + h + sigma * eps - z) / h, 3) # → 1.0 482round((mu + (sigma + h) * eps - z) / h, 3) # → 1.2 483``` 484 485```mermaid 486flowchart LR 487 subgraph A["Sampling directly: the gradient stops"] 488 direction LR 489 m1["μ, σ"] --> s1["draw z from N(μ, σ²)<br/>a dice roll"] --> d1["decoder"] --> l1["loss"] 490 l1 -. "no route back<br/>through a dice roll" .-> s1 491 end 492 subgraph B["Reparameterized: the gradient flows"] 493 direction LR 494 e2["ε from N(0, I)<br/>just another input"] --> s2["z = μ + σ·ε<br/>plain arithmetic"] 495 m2["μ, σ"] --> s2 --> d2["decoder"] --> l2["loss"] 496 l2 -. "∂z/∂μ = 1, ∂z/∂σ = ε" .-> m2 497 end 498``` 499 500**Reading it:** both rows produce the same kind of random code: a draw from a 501bell curve centred on μ with spread σ. In the top row the randomness sits 502*between* the encoder's outputs and the loss, and the dotted arrow of 503backpropagation has nowhere to go. In the bottom row the randomness comes in 504from the side as ε, an input like a pixel, and everything from μ and σ to the 505loss is ordinary arithmetic the chain rule can pass through. The encoder 506learns because of this one rearrangement. 507 508**Why it matters in practice.** The same trick lets gradients pass through 509random choices elsewhere too, for example in the noisy steps of diffusion 510models (`primer.ml.generative.diffusion`) and in reinforcement learning with 511continuous actions. It is a large part of why the VAE paper mattered. 512 513**In code:** `reparameterize` is the formula; in `VAE.loss_and_gradients` the lines under "Through z = μ + σ·ε" are the two gradients above, and a test checks them against finite differences. 514 515### The KL penalty: rent for every region 516 517**Everyday picture.** The rent from the everyday picture has a formal name: 518the **KL divergence** (Kullback-Leibler divergence) from the region to the 519standard bell curve, a measure of how different two distributions are that 520is 0 only when they match (`primer.ml.training_stages` uses it for 521distillation). For a bell-curve region and the standard bell curve, it has a 522short closed form: 523 524$$ 525\mathrm{KL}\bigl(\mathcal{N}(\mu, \sigma^2) \,\big\Vert\, \mathcal{N}(0, 1)\bigr) = \frac{1}{2} \sum_{j=1}^{k} \bigl( \mu_j^2 + \sigma_j^2 - \log \sigma_j^2 - 1 \bigr) 526$$ 527 528**Symbols** 529 530| Symbol | Meaning here | In the example | 531|---|---|---| 532| $\mathrm{KL}(q \,\Vert\, p)$ | how much distribution $q$ differs from distribution $p$; 0 when they match | | 533| $\mathcal{N}(\mu, \sigma^2)$ | the picture's region: a bell curve with centre μ and variance σ² | $\mathcal{N}(0.5, 0.25)$ | 534| $\mathcal{N}(0, 1)$ | the standard bell curve every region is pulled towards | | 535| $k$ | how many numbers in the code | 1 (the lesson's network uses 2) | 536| $j$ | a counter over the code numbers | | 537| $\mu_j^2$ | rent for sitting away from the centre: 0 at μ = 0, growing on both sides | 0.25 | 538| $\sigma_j^2 - \log \sigma_j^2$ | rent on size: smallest (exactly 1) when σ = 1, huge as σ shrinks to a pin, growing if it bloats | $0.25 + 1.386$ | 539| $-1$ | shifts the total so a perfect match costs exactly 0 | | 540| $\frac{1}{2}$ | a constant that falls out of the bell curve's formula | | 541 542**In words:** "for each code number, add the squared distance of its centre 543from 0, plus its variance, minus the log of its variance, minus 1; add those 544up over the code numbers and halve." 545 546**With the numbers:** for μ = 0.5 and σ = 0.5, 547½ (0.25 + 0.25 − (−1.386) − 1) = ½ × 0.886 = **0.443**. It splits into two 548rents: the centre alone costs ½ × 0.25 = 0.125, the shrunken size alone 549½ (0.25 + 1.386 − 1) = 0.318, and 0.125 + 0.318 = 0.443. A second code 550number that already has μ = 0 and σ = 1 adds ½ (0 + 1 − 0 − 1) = 0, so the 5512-number total is still 0.443. Shrink σ to a pin of 0.05 and the rent jumps 552to 2.62. 553 554**In Python:** 555 556```python 557import math 558def kl(mu, var): 559 return 0.5 * (mu ** 2 + var - math.log(var) - 1) 560# one code number with μ = 0.5, σ = 0.5 (σ² = 0.25) 561round(kl(0.5, 0.25), 3) # → 0.443 562# the rent for the centre alone, then for the size alone 563round(0.5 * 0.5 ** 2, 3) # → 0.125 564round(kl(0.0, 0.25), 3) # → 0.318 565# a code number that already is the standard normal pays nothing 566kl(0.0, 1.0) # → 0.0 567# Σ over j: two code numbers, their rents add 568round(kl(0.5, 0.25) + kl(0.0, 1.0), 3) # → 0.443 569# shrink σ to a pin (σ = 0.05) and the rent jumps 570round(kl(0.5, 0.05 ** 2), 2) # → 2.62 571``` 572 573 574 575**Reading it:** the left curve holds the spread at σ = 1 and moves the 576centre: a bowl with its bottom at μ = 0, so sitting at 0.5 costs 0.125 and 577sitting at 3 costs 4.5. The right curve holds the centre at 0 and changes 578the spread. It is 0 at σ = 1 (the dashed line), rises gently if the region 579bloats, and shoots up as σ heads towards 0. That steep wall on the left is 580what stops the encoder from shrinking every region to a pin and turning back 581into a plain autoencoder with holes. The two red dots add up to the worked 582example's 0.443. 583 584**In code:** `kl_to_standard_normal` is the formula, summed over the code numbers; a test checks it against a brute-force average over 400,000 samples. 585 586### The whole VAE loss 587 588$$ 589L = \bigl\lVert x - \hat{x} \bigr\rVert^2 + \beta \cdot \mathrm{KL}\bigl(q(z \mid x) \,\big\Vert\, \mathcal{N}(0, I)\bigr) 590$$ 591 592**Symbols** 593 594| Symbol | Meaning here | In the example | 595|---|---|---| 596| $x$ | the picture | (2, 3) from the line example | 597| $\hat{x}$ | the decoder's rebuild from the sampled code $z$ | (1.6, 3.2) | 598| $\lVert x - \hat{x} \rVert^2$ | the rebuild error, as in the plain autoencoder | 0.2 | 599| $q(z \mid x)$ | the region the encoder draws for picture $x$: read "q of z given x" | $\mathcal{N}(0.5, 0.25)$ | 600| $\beta$ | "beta": how much one unit of KL rent costs, in units of rebuild error | 0.3 | 601| KL(…) | the rent from the previous section | 0.443 | 602 603**In words:** "rebuild well, but pay β for every unit by which your regions 604differ from the standard bell curve." 605 606**With the numbers:** 0.2 + 0.3 × 0.443 = 0.2 + 0.133 = 0.333. 607 608**In Python:** 609 610```python 611rebuild, kl_value, beta = 0.2, 0.443, 0.3 612# L = ‖x − x̂‖² + β · KL 613round(rebuild + beta * kl_value, 3) # → 0.333 614``` 615 616Where does this come from? The VAE paper derives the loss (with β = 1) as the 617negative of the **evidence lower bound (ELBO)**: a quantity that is 618guaranteed never to exceed the log-probability the model gives to the real 619pictures, so pushing the bound up pushes that probability up. Choosing 620squared error amounts to assuming each pixel is the decoder's output plus 621bell-curve noise of spread s, and working that through gives β = 2s². This 622lesson uses β = 0.3, which assumes pixel noise of spread about 0.39. 623 624 625 626**Reading it:** the same 200 pictures as the plain autoencoder's map, now 627placed at their centres μ, each with its fuzzy region shaded around it 628(regions are drawn one σ wide). The dashed circles are radius 1 and 2 of the 629standard bell curve. Compare the axes with the plain autoencoder's: the codes 630now sit around 0 with a spread of about 1.1 in each direction, instead of 631wandering from −24 to 14. The strands are still there (the code still knows 632direction and offset), but they are packed side by side and the regions 633(σ about 0.14) overlap along each strand, so there is far less empty field 634left. 635 636**In code:** `VAE.loss_and_gradients` computes both terms and their gradients; `trained_vae` trains the lesson's VAE at a chosen β. 637 638## Sampling new strokes 639 640**Everyday picture.** Once the map is packed, throw a dart at the middle of 641it and you land in a neighbourhood, not a field. Generating is exactly that: 642throw away the encoder, draw a code from the standard bell curve, and decode. 643 644**Tiny example.** Draw the same code as before, z = (1.2, −0.9), and decode 645it with the VAE. Out comes a horizontal stroke through the centre of the 646picture, with squared distance about 0.06 from a real stroke in the training 647set. The plain autoencoder turned that same code into a smear 7.5 away. 648Over 500 draws the VAE's median distance is about 0.56 against the plain 649autoencoder's 3. 650 651```mermaid 652flowchart LR 653 N["draw z from N(0, I)<br/>2 random numbers"] --> D["decoder g"] --> X["a new picture<br/>never seen in training"] 654 ENC["encoder"] -. "not used when generating" .-> N 655``` 656 657**Reading it:** generating uses only the right half of the network. The 658encoder's job was done during training: it taught the decoder, through the 659KL rent, that codes live where the standard normal puts them. That is why the 660first box can draw from $\mathcal{N}(0, I)$ with confidence. The dotted arrow 661is a reminder that nothing flows from the encoder here. 662 663 664 665**Reading it:** both panels decode the *same* 24 random codes. The plain 666autoencoder (left) gives blobs, broken lines and a few strokes by luck. The 667VAE (right) gives strokes in all four directions at many offsets, plus a few 668soft blends of two directions. Those blends are the VAE's honest weak spot: 669with only 2 numbers to hold four directions, some codes sit on the border 670between two strands, and the decoder hedges between them. About a third of 671the VAE's samples still cross the junk line, most of them blends like these. 672 673 674 675**Reading it:** every tile is the VAE's decoding of one point on an even 676grid over the middle of the code map. Read along any row or column: the 677stroke changes a little at a time, sliding or turning, with no sudden jumps 678to junk. Whole regions of the map belong to one direction, and the borders 679between them are where the blends live. This is what a **latent space** 680means: a space of codes where position means something. That smoothness is 681what lets you edit or explore a picture by moving its code. 682 683### Interpolation: walking between two codes 684 685**Everyday picture.** A morph between two faces in a film: every in-between 686frame should look like a face, not a double exposure. 687 688**Tiny example.** Encode two pictures to codes $z_a$ and $z_b$, step along 689the straight line between them, and decode every stop: 690 691$$ 692z_t = (1 - t)\, z_a + t\, z_b, \qquad 0 \le t \le 1 693$$ 694 695**Symbols** 696 697| Symbol | Meaning here | In the example | 698|---|---|---| 699| $z_a, z_b$ | the codes of the two pictures at the ends of the walk | $(-1, 0.5)$ and $(1, 1.5)$ | 700| $t$ | how far along the walk: 0 at $z_a$, 1 at $z_b$ | 0.25 | 701| $z_t$ | the code at that point of the walk | $(-0.5, 0.75)$ | 702 703**In words:** "take a share $1 - t$ of the first code and a share $t$ of the 704second, and add them." 705 706**With the numbers:** at $t = 0.25$, $z_t = 0.75 \times (-1, 0.5) + 0.25 707\times (1, 1.5) = (-0.75 + 0.25,\ 0.375 + 0.375) = (-0.5, 0.75)$. The walk in 708the figure uses nine stops, $t = 0, 1/8, 2/8, \ldots, 1$. 709 710**In Python:** 711 712```python 713z_a, z_b = [-1.0, 0.5], [1.0, 1.5] 714# z_t = (1 − t) z_a + t z_b, a quarter of the way along 715t = 0.25 716[(1 - t) * a + t * b for a, b in zip(z_a, z_b)] # → [-0.5, 0.75] 717# the nine stops of a walk 718[i / 8 for i in range(9)] # → [0.0, 0.125, 0.25, 0.375, 0.5, 0.625, 0.75, 0.875, 1.0] 719``` 720 721 722 723**Reading it:** both rows walk between the same two vertical strokes, one 724left of centre and one right of it. The plain autoencoder's walk (top) 725leaves its strand and crosses a hole: the middle frames are a smeared 726cross, 3.88 away from any real stroke. The VAE's walk (bottom) takes a 727surprising route, through diagonal strokes, because in its 2-number map the 728diagonal strand sits between those two points. But every stop is a 729believable stroke, never more than 0.70 from a real one. Averaged over 60 730random pairs, the VAE's biggest jump from one frame to the next is also 731smaller (about 1.6 against 2.4). 732 733**Why it matters in practice.** Smooth latent spaces are what made "move 734this slider to add a smile" demos possible, and interpolation is still a 735standard sanity check for any learned code: if the walk is full of junk, the 736space has holes. 737 738**In code:** `generate` draws and decodes; `interpolate` encodes two pictures and decodes the walk between them; `largest_step` measures its biggest jump. 739 740## The trade-off: β, blur and collapse 741 742**Everyday picture.** Go back to the rent. Set it too low and every region 743shrinks to a pin far from the centre: rebuilds are sharp, but the holes come 744back. Set it too high and every region moves to the centre and swells to 745the full bell curve. Then every picture's region is the same region, the 746code carries no information at all, and the decoder can do nothing better 747than draw the average of all the strokes: a grey smudge. That failure is 748called **posterior collapse**. 749 750**Tiny example.** The lesson trains six VAEs, with β from 0.01 to 3: 751 752| β | rebuild error | KL (information in the code) | median sample distance to a real stroke | 753|---|---|---|---| 754| 0.01 | 0.16 | 10.3 | 2.40 | 755| 0.1 | 0.23 | 5.7 | 0.87 | 756| 0.3 | 0.31 | 4.4 | 0.56 | 757| 1 | 0.92 | 2.9 | 0.81 | 758| 3 | 6.08 | 0.0 | 4.73 | 759 760(`beta_sweep` produces this table, including β = 0.03, and `demo` prints it.) 761 762The samples are best in the middle. At β = 3 the KL is zero: the regions are 763the standard bell curve itself, and the rebuild error of 6.08 is what you get 764by drawing the average stroke every time. 765 766### Why VAE pictures are blurry 767 768Even at a good β the regions overlap, so one code can stand for several 769different pictures. What should the decoder draw for it? Take a single pixel 770that is black (1) in one of those pictures and white (0) in the other, equally 771likely. Squared error has a clear answer: 772 773$$ 774\underset{c}{\arg\min}\ \mathbb{E}\bigl[(x - c)^2\bigr] = \mathbb{E}[x] 775$$ 776 777**Symbols** 778 779| Symbol | Meaning here | In the example | 780|---|---|---| 781| $x$ | the pixel's true value, which could be either picture's | 0 or 1, equally likely | 782| $c$ | the decoder's guess for that pixel | 0, 0.5 or 1 | 783| $(x - c)^2$ | the squared error of the guess | | 784| $\mathbb{E}[\ldots]$ | **expected value**: the average over the possibilities, each weighted by its chance | | 785| $\arg\min_c$ | "the $c$ that makes what follows as small as possible" | | 786 787**In words:** "the guess with the smallest average squared error is the 788average of the possibilities." 789 790**With the numbers:** guess black (1): error 1 half the time, 0 otherwise, 791average 0.5. Guess white (0): also 0.5. Guess grey (0.5): error 0.25 either 792way, average **0.25**. Grey wins, and grey is the average of 0 and 1. 793 794**In Python:** 795 796```python 797outcomes = [0.0, 1.0] 798def expected_error(c): 799 return sum((x - c) ** 2 for x in outcomes) / len(outcomes) 800# E[(x − c)²] for three guesses: white, grey, black 801[expected_error(c) for c in (0.0, 0.5, 1.0)] # → [0.5, 0.25, 0.5] 802# E[x]: the average outcome, which is the winning guess 803sum(outcomes) / len(outcomes) # → 0.5 804``` 805 806Across a whole picture, "average the pictures this code might mean" is a 807blur. The more the regions overlap, the more pictures each code must stand 808for, and the blurrier the output. 809 810 811 812**Reading it:** on the left, β grows along a logarithmic axis. The blue 813rebuild error stays low, then climbs steeply past β = 1. The green KL, the 814information the code carries, falls steadily and hits zero at β = 3. The red 815line, how far random samples are from real strokes, is a U: high at small β 816(holes), lowest near β = 0.3, high again at large β (blur and collapse). 817On the right are eight samples at each β. The top rows have sharp ink in 818junk shapes. The middle rows are clean strokes. The β = 1 row is softer. 819The bottom row is eight copies of the same grey smudge: the collapse. 820 821**Why it matters in practice.** β is a real dial, and the name β-VAE 822(Higgins and colleagues, 2017) comes from turning it up to get codes whose 823numbers line up with separate factors, like direction and offset here. The 824blur is the VAE's best-known weakness. Modern systems fix it by adding a 825second judge of realism to the loss (an adversarial loss, the trick behind 826`primer.ml.generative.gans`, as in VQGAN), or by letting a stronger generator 827do the generating while the autoencoder only compresses. 828 829**In code:** `beta_sweep` trains and measures one VAE per β; `expected_squared_error` is the grey-pixel calculation. 830 831## Where autoencoders live today 832 833**Everyday picture.** The autoencoder became the zip format for images and 834sound that other generative models work inside. 835 836**Tiny example.** Stable Diffusion's autoencoder turns a 512 × 512 colour 837image (512 × 512 × 3 = 786,432 numbers) into a 64 × 64 × 4 grid of codes 838(16,384 numbers): 48 times fewer. The diffusion model does all its work on 839that small grid, and the decoder paints full-size pixels only at the very 840end. Its KL weight is tiny, so it is mostly a compressor, with just enough 841rent to keep the codes well-behaved. 842 843A **VQ-VAE** (vector-quantized VAE) goes one step further: it snaps each code 844vector to the nearest entry of a learned **codebook**, so an image becomes a 845grid of whole numbers, the entries' positions in the codebook. Those numbers 846are **tokens**, and a transformer can read and write them exactly as it 847reads and writes words. Neural audio codecs do the same for sound. 848 849```mermaid 850flowchart LR 851 subgraph LD["Latent diffusion"] 852 direction LR 853 I1["image<br/>512 × 512 × 3"] --> E1["VAE encoder"] --> L1["latent<br/>64 × 64 × 4"] 854 L1 --> DF["diffusion model<br/>works here"] --> D1["VAE decoder"] --> O1["image"] 855 end 856 subgraph TK["Tokenizer for a transformer"] 857 direction LR 858 I2["image or audio"] --> E2["encoder"] --> Q["snap each vector to<br/>nearest codebook entry"] 859 Q --> T["grid of token ids"] --> TR["transformer"] 860 TR --> D2["decoder"] --> O2["image or audio"] 861 end 862``` 863 864**Reading it:** in the top row the autoencoder is the outer shell: its 865encoder shrinks the image 48-fold on the way in, its decoder restores it on 866the way out, and the expensive generator in the middle never touches a 867pixel. Training and sampling run far faster on 16,384 numbers than on 868786,432, which is what made high-resolution diffusion affordable (see 869`primer.ml.generative.diffusion`). In the bottom row the snapping step turns 870continuous codes into token ids, which is how images and audio can enter 871and leave a language-model-style transformer (see 872`primer.ml.generative.multimodal`). 873 874**Why it matters in practice.** When a model "generates an image", there is 875very often an autoencoder at both ends of the pipeline. Its quality sets a 876ceiling: whatever detail the decoder cannot rebuild, no generator working in 877its latent space can produce. 878 879## In 20 seconds 880 881- **Autoencoder:** an encoder squeezes data through a narrow code (the 882 bottleneck) and a decoder rebuilds it; training minimises the rebuild 883 error. With no bends it learns exactly what PCA learns; with bends it can 884 follow curved data far better. 885- **No generation from a plain autoencoder:** its codes land wherever 886 training put them, with holes in between, so random codes decode to junk. 887- **VAE:** the encoder outputs a region (μ, σ); the reparameterization trick 888 z = μ + σ·ε samples from it in a way gradients can pass through; a KL 889 penalty pulls every region towards the standard normal, so drawing 890 z ~ N(0, I) and decoding generates new data. 891- **The trade-off:** β weighs the KL. Too small brings back holes, too large 892 blurs and finally collapses; squared error averages over overlapping 893 possibilities, which is why VAE output is blurry. 894- **Today:** VAEs compress images for latent diffusion, and VQ-VAEs turn 895 images and audio into tokens for transformers. 896 897## Self-test questions 898 899**What is the bottleneck for? What goes wrong if the code is as wide as the 900input?** 901It forces the network to keep only what matters: with 2 numbers for 64 902pixels, it must find the few facts that actually vary. If the code is as 903wide as the input, the network can learn to copy the pixels straight through, 904rebuild perfectly and learn nothing useful, unless something else (noise on 905the input, a penalty on the code) stops it. 906 907**Why can't you generate new pictures by decoding random codes from a plain 908autoencoder?** 909Its loss only ever sees the codes of real pictures, so it says nothing about 910where codes should live or what lies between them. The codes end up in an 911arbitrary range with holes between clusters, and a random code usually lands 912in a hole, where the decoder's output is junk. 913 914**What problem does the reparameterization trick solve?** 915Backpropagation needs every step between the weights and the loss to be 916something you can differentiate, and a random draw is not. Writing the draw 917as z = μ + σ·ε with the noise ε as a separate input turns sampling into 918arithmetic, so gradients reach μ and σ, and through them the encoder. 919 920**What do the two terms of the VAE loss each want, and why do you need 921both?** 922The rebuild term wants regions small and far apart so every picture is 923decoded precisely. The KL term wants every region to be the standard bell 924curve. Without KL you get a plain autoencoder with holes; without the 925rebuild term every region collapses onto the bell curve and the code carries 926nothing. The balance gives a packed, smooth code space that still tells 927pictures apart. 928 929**Work out the KL penalty for one code number with μ = 0 and σ = 2.** 930½ (0 + 4 − log 4 − 1) = ½ (3 − 1.386) = 0.807. A region that is too big pays 931rent too, though less steeply than one that is too small. 932 933**Why are VAE samples blurry?** 934Regions overlap, so one code can stand for several pictures, and under 935squared error the best single answer is their average. Averaged pictures are 936blurred pictures. Raising β increases the overlap and the blur. 937 938**What is posterior collapse?** 939When the KL rent outweighs what the code saves in rebuild error, the encoder 940makes every region the standard bell curve, the code carries no information, 941and the decoder outputs the same average picture for every code. In this 942lesson that happens at β = 3. 943 944**Why does latent diffusion run inside an autoencoder's code space instead of 945on pixels?** 946The code is many times smaller (48 times for Stable Diffusion's 512 × 512 947images) and keeps what matters to the eye, so the expensive, many-step 948generator is far cheaper to train and run. The decoder turns the result back 949into full-size pixels once, at the end. 950 951## The papers behind this lesson 952 953- **Hinton & Salakhutdinov, *Reducing the Dimensionality of Data with Neural 954 Networks* (Science, 2006)**: https://doi.org/10.1126/science.1127647. 955 Showed that deep autoencoders, once they could be trained, compress data 956 far better than PCA. 957- **Kingma & Welling, *Auto-Encoding Variational Bayes* (2013)**: 958 https://arxiv.org/abs/1312.6114. Introduced the variational autoencoder, 959 the reparameterization trick and the ELBO loss with its closed-form 960 Gaussian KL. 961 [Annotated companion](../../../papers/auto-encoding-variational-bayes.html) 962- **Rezende, Mohamed & Wierstra, *Stochastic Backpropagation and Approximate 963 Inference in Deep Generative Models* (2014)**: 964 https://arxiv.org/abs/1401.4082. Developed the same idea independently at 965 the same time, showing how to backpropagate through random sampling. 966- **Burgess et al., *Understanding disentangling in β-VAE* (2018)**: 967 https://arxiv.org/abs/1804.03599. Explains why turning up β pushes the 968 code's numbers to line up with separate factors of the data, and what it 969 costs in rebuild quality. 970- **van den Oord, Vinyals & Kavukcuoglu, *Neural Discrete Representation 971 Learning* (2017)**: https://arxiv.org/abs/1711.00937. Introduced the 972 VQ-VAE, which snaps codes to a learned codebook and so turns images and 973 audio into tokens. 974 [Annotated companion](../../../papers/vq-vae.html) 975- **Rombach et al., *High-Resolution Image Synthesis with Latent Diffusion 976 Models* (2021)**: https://arxiv.org/abs/2112.10752. Ran diffusion inside a 977 lightly regularized autoencoder's code space, the design behind Stable 978 Diffusion. 979 980## Further reading 981 982- Kingma & Welling, *An Introduction to Variational Autoencoders* (2019): https://arxiv.org/abs/1906.02691 983- Carl Doersch, *Tutorial on Variational Autoencoders* (2016): https://arxiv.org/abs/1606.05908 984- Goodfellow, Bengio & Courville, *Deep Learning*, chapter 14, Autoencoders: https://www.deeplearningbook.org/contents/autoencoders.html 985- Esser, Rombach & Ommer, *Taming Transformers for High-Resolution Image Synthesis* (VQGAN, 2020): https://arxiv.org/abs/2012.09841 986- PyTorch's VAE example: https://github.com/pytorch/examples/tree/main/vae 987""" 988 989from __future__ import annotations 990 991from functools import lru_cache 992 993import numpy as np 994 995from primer._show import banner, say, table, takeaway 996from primer.ml.optimizers import Adam 997 998# --------------------------------------------------------------------------- 999# 1. The toy data: 8 × 8 pictures of pen strokes 1000# --------------------------------------------------------------------------- 1001 1002SIDE = 8 # images are SIDE × SIDE pixels, flattened to 64 numbers 1003STROKE_KINDS = ("horizontal", "vertical", "diagonal down", "diagonal up") 1004_STROKE_ANGLES = (0.0, np.pi / 2, np.pi / 4, -np.pi / 4) # radians, one per kind above 1005STROKE_WIDTH = 0.6 # pixels: how quickly the ink fades either side of the line 1006 1007# A real stroke carries about 8.8 units of squared ink (the sum of its squared 1008# pixels). An image whose squared distance to every training stroke is above 1009# 1.0 is more than about a ninth of a stroke away from anything real: junk. 1010JUNK_DISTANCE = 1.0 1011 1012 1013def make_strokes(n: int = 200, seed: int = 0) -> tuple[np.ndarray, np.ndarray, np.ndarray]: 1014 """n pictures of one soft pen stroke each: (images (n, 64), kinds (n,), offsets (n,)). 1015 1016 Two hidden facts make each picture: which of the four directions the 1017 stroke runs (kinds cycle 0, 1, 2, 3) and how far it sits from the centre, 1018 a continuous offset in [-2.5, 2.5] pixels. A good 2-number code has to 1019 rediscover both from the 64 pixels alone. 1020 """ 1021 rng = np.random.default_rng(seed) 1022 kinds = np.arange(n) % len(STROKE_KINDS) 1023 offsets = rng.uniform(-2.5, 2.5, n) 1024 # Pixel centres measured from the middle of the image, so offset 0 passes through the centre. 1025 rows, cols = np.mgrid[0:SIDE, 0:SIDE] - (SIDE - 1) / 2 1026 angle = np.array(_STROKE_ANGLES)[kinds] 1027 # Signed distance from each pixel to the line: project the pixel onto the 1028 # line's perpendicular direction, then subtract how far the line is shifted. 1029 across_x, across_y = -np.sin(angle), np.cos(angle) 1030 dist = across_x[:, None, None] * cols + across_y[:, None, None] * rows - offsets[:, None, None] 1031 # Ink fades like a bell curve with distance: 1.0 on the line, about 0.25 one pixel away. 1032 images = np.exp(-(dist**2) / (2 * STROKE_WIDTH**2)) 1033 return images.reshape(n, SIDE * SIDE), kinds, offsets 1034 1035 1036# --------------------------------------------------------------------------- 1037# 2. The smallest autoencoder: points on a line, a 1-number code 1038# --------------------------------------------------------------------------- 1039 1040 1041def worked_example_line() -> dict: 1042 """Squeeze 2-D points to one number and back, for points near y = 2x. 1043 1044 Encoder: code = u · x, with u = (1, 2)/√5, the line's direction scaled to 1045 length 1. Decoder: rebuilt = code · u. A point on the line survives the 1046 trip exactly; a point off it comes back as its closest point on the line. 1047 """ 1048 u = np.array([1.0, 2.0]) / np.sqrt(5) 1049 on, off = np.array([1.0, 2.0]), np.array([2.0, 3.0]) 1050 on_code, off_code = float(u @ on), float(u @ off) 1051 off_rebuilt = off_code * u 1052 return dict( 1053 direction=u, 1054 on_line_code=on_code, 1055 on_line_rebuilt=on_code * u, 1056 off_line_code=off_code, 1057 off_line_rebuilt=off_rebuilt, 1058 off_line_error=float(np.sum((off - off_rebuilt) ** 2)), 1059 ) 1060 1061 1062def pca_reconstruction(X: np.ndarray, k: int) -> np.ndarray: 1063 """Project X onto its top-k principal directions and back (see `primer.ml.embeddings.clustering.pca`). 1064 1065 PCA is the best any *flat* k-number code can do under squared error. 1066 """ 1067 mean = X.mean(axis=0) 1068 _, _, Vt = np.linalg.svd(X - mean, full_matrices=False) 1069 return (X - mean) @ Vt[:k].T @ Vt[:k] + mean 1070 1071 1072def train_linear_autoencoder(X: np.ndarray, n_code: int = 1, steps: int = 2000, lr: float = 0.02, seed: int = 0): 1073 """An autoencoder with no bends: code = x W_enc, rebuilt = code W_dec. Returns (W_enc, W_dec). 1074 1075 Trained by plain gradient descent on the squared error of centred data. 1076 With nothing nonlinear in it, the best it can do is PCA's answer: it ends 1077 up spanning the same directions as the top principal components. 1078 """ 1079 rng = np.random.default_rng(seed) 1080 Xc = X - X.mean(axis=0) 1081 n, d = Xc.shape 1082 W_enc = rng.normal(0, 0.1, (d, n_code)) # (d, k) 1083 W_dec = rng.normal(0, 0.1, (n_code, d)) # (k, d) 1084 for _ in range(steps): 1085 Z = Xc @ W_enc # (n, k) codes 1086 R = Z @ W_dec - Xc # (n, d) what the rebuild got wrong 1087 # Loss = mean over points of the summed squared error; these are its exact gradients. 1088 g_dec = 2 / n * Z.T @ R 1089 g_enc = 2 / n * Xc.T @ (R @ W_dec.T) 1090 W_enc -= lr * g_enc 1091 W_dec -= lr * g_dec 1092 return W_enc, W_dec 1093 1094 1095# --------------------------------------------------------------------------- 1096# 3. A nonlinear autoencoder, and its variational cousin, with hand-written backprop 1097# --------------------------------------------------------------------------- 1098 1099 1100def sigmoid(z: np.ndarray) -> np.ndarray: 1101 # Squashes any number into (0, 1): the decoder's pixels must be valid brightnesses. 1102 return 1 / (1 + np.exp(-z)) 1103 1104 1105class Autoencoder: 1106 """64 pixels -> 32 hidden -> a code of `n_code` numbers -> 32 hidden -> 64 pixels. 1107 1108 The encoder and the decoder are each a two-layer network like the one in 1109 `primer.ml.neural_net`. Training asks only one thing: rebuild the input. 1110 """ 1111 1112 def __init__(self, n_in: int = SIDE * SIDE, n_hidden: int = 32, n_code: int = 2, seed: int = 0): 1113 self.n_in, self.n_hidden, self.n_code = n_in, n_hidden, n_code 1114 rng = np.random.default_rng(seed) 1115 head = self._head_width() 1116 # Variance 1/fan_in keeps tanh out of its flat regions at the start (see primer.ml.deep_nets). 1117 self.params = { 1118 "W1": rng.normal(0, 1 / np.sqrt(n_in), (n_in, n_hidden)), # encoder 1119 "b1": np.zeros(n_hidden), 1120 "W2": rng.normal(0, 1 / np.sqrt(n_hidden), (n_hidden, head)), 1121 "b2": np.zeros(head), 1122 "W3": rng.normal(0, 1 / np.sqrt(n_code), (n_code, n_hidden)), # decoder 1123 "b3": np.zeros(n_hidden), 1124 "W4": rng.normal(0, 1 / np.sqrt(n_hidden), (n_hidden, n_in)), 1125 "b4": np.zeros(n_in), 1126 } 1127 1128 def _head_width(self) -> int: 1129 return self.n_code # the encoder's last layer outputs the code itself 1130 1131 def _encoder(self, X: np.ndarray) -> tuple[np.ndarray, np.ndarray]: 1132 p = self.params 1133 h = np.tanh(X @ p["W1"] + p["b1"]) # (n, hidden) 1134 return h, h @ p["W2"] + p["b2"] # (n, head): no squashing, a code may be any number 1135 1136 def _decoder(self, Z: np.ndarray) -> tuple[np.ndarray, np.ndarray]: 1137 p = self.params 1138 g = np.tanh(Z @ p["W3"] + p["b3"]) # (n, hidden) 1139 return g, sigmoid(g @ p["W4"] + p["b4"]) # (n, 64) pixels in (0, 1) 1140 1141 def encode(self, X: np.ndarray) -> np.ndarray: 1142 """Images (n, 64) -> codes (n, n_code).""" 1143 return self._encoder(X)[1] 1144 1145 def decode(self, Z: np.ndarray) -> np.ndarray: 1146 """Codes (n, n_code) -> images (n, 64). Works on any code, seen in training or not.""" 1147 return self._decoder(Z)[1] 1148 1149 def reconstruct(self, X: np.ndarray) -> np.ndarray: 1150 return self.decode(self.encode(X)) 1151 1152 def _backward_decoder(self, Z, g, X_hat, X) -> tuple[dict, np.ndarray]: 1153 """Gradients of the mean summed squared error for the decoder, and the error signal at the code.""" 1154 p, n = self.params, len(X) 1155 d_out = 2 * (X_hat - X) / n * X_hat * (1 - X_hat) # through the squared error, then the sigmoid 1156 d_g = d_out @ p["W4"].T * (1 - g**2) # back through W4, then tanh' = 1 - tanh² 1157 grads = {"W4": g.T @ d_out, "b4": d_out.sum(axis=0), "W3": Z.T @ d_g, "b3": d_g.sum(axis=0)} 1158 return grads, d_g @ p["W3"].T # (n, n_code): how the loss changes with each code number 1159 1160 def _backward_encoder(self, X, h, d_head) -> dict: 1161 p = self.params 1162 d_h = d_head @ p["W2"].T * (1 - h**2) 1163 return {"W2": h.T @ d_head, "b2": d_head.sum(axis=0), "W1": X.T @ d_h, "b1": d_h.sum(axis=0)} 1164 1165 def loss_and_gradients(self, X: np.ndarray, noise: np.ndarray | None = None) -> tuple[dict, dict]: 1166 """Reconstruction loss (squared error summed over pixels, averaged over images) and every gradient. 1167 1168 `noise` is ignored; it is accepted so `train` can drive this and `VAE` the same way. 1169 """ 1170 h, Z = self._encoder(X) 1171 g, X_hat = self._decoder(Z) 1172 rec = float(np.sum((X_hat - X) ** 2) / len(X)) 1173 grads, d_Z = self._backward_decoder(Z, g, X_hat, X) 1174 grads |= self._backward_encoder(X, h, d_Z) 1175 return {"loss": rec, "reconstruction": rec, "kl": 0.0}, grads 1176 1177 1178class VAE(Autoencoder): 1179 """A variational autoencoder: the encoder outputs a fuzzy region, not a point. 1180 1181 For each image the encoder gives a mean μ and a log-variance log σ² per 1182 code number. Training draws a code from that region with the 1183 reparameterization trick and adds `beta` times the KL penalty that pulls 1184 every region towards the standard normal distribution. 1185 """ 1186 1187 def __init__(self, n_in: int = SIDE * SIDE, n_hidden: int = 32, n_code: int = 2, beta: float = 0.3, seed: int = 0): 1188 self.beta = beta 1189 super().__init__(n_in, n_hidden, n_code, seed) 1190 1191 def _head_width(self) -> int: 1192 return 2 * self.n_code # μ and log σ² for every code number 1193 1194 def encode_distribution(self, X: np.ndarray) -> tuple[np.ndarray, np.ndarray]: 1195 """Images -> (μ, log σ²), each (n, n_code).""" 1196 head = self._encoder(X)[1] 1197 return head[:, : self.n_code], head[:, self.n_code :] 1198 1199 def encode(self, X: np.ndarray) -> np.ndarray: 1200 """The centre μ of each image's region: the single best code to decode it from.""" 1201 return self.encode_distribution(X)[0] 1202 1203 def loss_and_gradients(self, X: np.ndarray, noise: np.ndarray | None = None) -> tuple[dict, dict]: 1204 """Reconstruction + β·KL, with gradients through the sampling step. 1205 1206 `noise` is ε, shape (n, n_code), drawn from the standard normal by the 1207 caller. Passing it in keeps the step deterministic, which is what lets 1208 the tests check every gradient against finite differences. 1209 """ 1210 n = len(X) 1211 noise = np.zeros((n, self.n_code)) if noise is None else noise 1212 h, head = self._encoder(X) 1213 mu, logvar = head[:, : self.n_code], head[:, self.n_code :] 1214 sigma = np.exp(0.5 * logvar) 1215 Z = reparameterize(mu, logvar, noise) # μ + σ·ε: randomness enters only through ε 1216 g, X_hat = self._decoder(Z) 1217 rec = float(np.sum((X_hat - X) ** 2) / n) 1218 kl = float(np.sum(kl_to_standard_normal(mu, logvar)) / n) 1219 grads, d_Z = self._backward_decoder(Z, g, X_hat, X) 1220 # Through z = μ + σ·ε: ∂z/∂μ = 1 and ∂z/∂σ = ε, with ∂σ/∂(log σ²) = σ/2. 1221 # The KL term adds its own pull: ∂KL/∂μ = μ and ∂KL/∂(log σ²) = (σ² - 1)/2. 1222 d_mu = d_Z + self.beta * mu / n 1223 d_logvar = d_Z * noise * sigma / 2 + self.beta * (sigma**2 - 1) / (2 * n) 1224 grads |= self._backward_encoder(X, h, np.hstack([d_mu, d_logvar])) 1225 return {"loss": rec + self.beta * kl, "reconstruction": rec, "kl": kl}, grads 1226 1227 1228def reparameterize(mu: np.ndarray, logvar: np.ndarray, eps: np.ndarray) -> np.ndarray: 1229 """z = μ + σ·ε with σ = exp(½ log σ²): a draw from N(μ, σ²) written as a plain sum. 1230 1231 The encoder predicts log σ² rather than σ because a log can be any number, 1232 positive or negative, while σ must stay positive; exp takes care of that. 1233 """ 1234 return mu + np.exp(0.5 * logvar) * eps 1235 1236 1237def kl_to_standard_normal(mu: np.ndarray, logvar: np.ndarray) -> np.ndarray: 1238 """KL(N(μ, σ²) ‖ N(0, 1)) summed over the last axis: ½ Σ (μ² + σ² - log σ² - 1).""" 1239 return 0.5 * np.sum(mu**2 + np.exp(logvar) - logvar - 1, axis=-1) 1240 1241 1242def train(model: Autoencoder, X: np.ndarray, steps: int = 1500, lr: float = 0.01, seed: int = 0) -> list[dict]: 1243 """Full-batch training with Adam (see `primer.ml.optimizers`). Returns the losses at every step.""" 1244 rng = np.random.default_rng(seed) 1245 optimizers = {name: Adam(lr=lr) for name in model.params} 1246 history = [] 1247 for _ in range(steps): 1248 # Fresh ε every step, so each image's code lands somewhere new in its region. 1249 noise = rng.standard_normal((len(X), model.n_code)) 1250 losses, grads = model.loss_and_gradients(X, noise) 1251 for name in model.params: 1252 model.params[name] = optimizers[name].step(model.params[name], grads[name]) 1253 history.append(losses) 1254 return history 1255 1256 1257@lru_cache(maxsize=None) 1258def trained_autoencoder(steps: int = 1500) -> Autoencoder: 1259 """The lesson's plain autoencoder, trained once on `make_strokes` and reused.""" 1260 model = Autoencoder() 1261 train(model, make_strokes()[0], steps=steps) 1262 return model 1263 1264 1265@lru_cache(maxsize=None) 1266def trained_vae(beta: float = 0.3, steps: int = 1500) -> VAE: 1267 """The lesson's VAE at a given β, trained once on `make_strokes` and reused.""" 1268 model = VAE(beta=beta) 1269 train(model, make_strokes()[0], steps=steps) 1270 return model 1271 1272 1273# --------------------------------------------------------------------------- 1274# 4. Measuring: rebuilds, junk, samples and walks 1275# --------------------------------------------------------------------------- 1276 1277 1278def reconstruction_error(model: Autoencoder, X: np.ndarray) -> float: 1279 """Squared error summed over the 64 pixels, averaged over images.""" 1280 return float(np.mean(np.sum((model.reconstruct(X) - X) ** 2, axis=1))) 1281 1282 1283def distance_to_data(images: np.ndarray, X: np.ndarray) -> np.ndarray: 1284 """For each image, the squared distance to the nearest training image. Large means junk.""" 1285 # ‖a - b‖² = ‖a‖² + ‖b‖² - 2 a·b, for every pair at once: (n_images, n_train). 1286 d = np.sum(images**2, axis=1)[:, None] + np.sum(X**2, axis=1)[None, :] - 2 * images @ X.T 1287 return np.maximum(d.min(axis=1), 0.0) 1288 1289 1290def generate(model: Autoencoder, n: int, seed: int = 0) -> np.ndarray: 1291 """Draw n codes from the standard normal and decode them: brand-new images.""" 1292 Z = np.random.default_rng(seed).standard_normal((n, model.n_code)) 1293 return model.decode(Z) 1294 1295 1296def codes_in_box(codes: np.ndarray, n: int, seed: int = 0) -> np.ndarray: 1297 """n codes drawn evenly from the smallest box that holds every given code.""" 1298 rng = np.random.default_rng(seed) 1299 return rng.uniform(codes.min(axis=0), codes.max(axis=0), (n, codes.shape[1])) 1300 1301 1302def interpolate(model: Autoencoder, x_a: np.ndarray, x_b: np.ndarray, steps: int = 9) -> np.ndarray: 1303 """Encode two images, walk a straight line between their codes, decode every stop: (steps, 64).""" 1304 z_a, z_b = model.encode(np.stack([x_a, x_b])) 1305 t = np.linspace(0, 1, steps)[:, None] # 0 at image a, 1 at image b 1306 return model.decode((1 - t) * z_a + t * z_b) 1307 1308 1309def largest_step(frames: np.ndarray) -> float: 1310 """The biggest change (Euclidean distance in pixels) between consecutive frames of a walk.""" 1311 return float(np.max(np.linalg.norm(np.diff(frames, axis=0), axis=1))) 1312 1313 1314def expected_squared_error(guess: float, outcomes) -> float: 1315 """Average squared error of one guess against equally likely outcomes.""" 1316 return float(np.mean([(o - guess) ** 2 for o in outcomes])) 1317 1318 1319def beta_sweep(betas=(0.01, 0.03, 0.1, 0.3, 1.0, 3.0)) -> list[dict]: 1320 """Train a VAE at each β and measure rebuild quality, KL, and how real its samples look.""" 1321 X = make_strokes()[0] 1322 rows = [] 1323 for beta in betas: 1324 model = trained_vae(beta) 1325 samples = generate(model, 500, seed=5) 1326 rows.append( 1327 dict( 1328 beta=beta, 1329 reconstruction=reconstruction_error(model, X), 1330 kl=float(np.mean(kl_to_standard_normal(*model.encode_distribution(X)))), 1331 sample_distance=float(np.median(distance_to_data(samples, X))), 1332 sample_peak=float(samples.max(axis=1).mean()), 1333 ) 1334 ) 1335 return rows 1336 1337 1338# --------------------------------------------------------------------------- 1339# 5. Figures (rendered into the HTML docs by `make figures`) 1340# --------------------------------------------------------------------------- 1341 1342_KIND_COLOURS = ("#2563eb", "#dc2626", "#16a34a", "#9333ea") # one per stroke kind 1343_JUNK, _REAL, _MUTED = "#dc2626", "#2563eb", "#9ca3af" 1344_WALK = (1, 45) # two vertical strokes, 3.1 px apart, used for the walk figure and the demo 1345 1346 1347def _tile(ax, images: np.ndarray, cols: int, title: str | None = None) -> None: 1348 """Draw 64-pixel images as one mosaic of 8 × 8 tiles, dark ink on white, with thin gaps.""" 1349 rows = int(np.ceil(len(images) / cols)) 1350 mosaic = np.full((rows * (SIDE + 1) - 1, cols * (SIDE + 1) - 1), np.nan) 1351 for i, img in enumerate(images): 1352 r, c = divmod(i, cols) 1353 mosaic[r * (SIDE + 1) : r * (SIDE + 1) + SIDE, c * (SIDE + 1) : c * (SIDE + 1) + SIDE] = img.reshape(SIDE, SIDE) 1354 ax.imshow(np.ma.masked_invalid(mosaic), cmap="Greys", vmin=0, vmax=1, interpolation="nearest") 1355 ax.set_facecolor("#e5e7eb") # the gaps between tiles 1356 ax.set_xticks([]) 1357 ax.set_yticks([]) 1358 if title: 1359 ax.set_title(title, fontsize=10) 1360 1361 1362def _pick_examples(kinds: np.ndarray, offsets: np.ndarray, per_kind: int = 2) -> list[int]: 1363 """Indices of a few strokes of every kind, spread across offsets, for display.""" 1364 picks = [] 1365 for k in range(len(STROKE_KINDS)): 1366 idx = np.where(kinds == k)[0] 1367 idx = idx[np.argsort(offsets[idx])] 1368 picks += [int(idx[int(q * (len(idx) - 1))]) for q in np.linspace(0.15, 0.85, per_kind)] 1369 return picks 1370 1371 1372def figures() -> dict: 1373 """Plot this lesson's data. matplotlib is imported here, and only here, 1374 so the lesson itself needs nothing beyond NumPy.""" 1375 import matplotlib 1376 1377 matplotlib.use("Agg") 1378 import matplotlib.pyplot as plt 1379 from matplotlib.patches import Circle, Ellipse, Rectangle 1380 1381 X, kinds, offsets = make_strokes() 1382 ae, vae = trained_autoencoder(), trained_vae() 1383 figs = {} 1384 1385 # --- 1. Rebuilds: autoencoder vs PCA, both with 2 numbers ---------------- 1386 picks = _pick_examples(kinds, offsets, per_kind=2) 1387 fig, axes = plt.subplots(3, 1, figsize=(7, 3.6)) 1388 _tile(axes[0], X[picks], 8, "the original pictures (64 pixels each)") 1389 _tile(axes[1], ae.reconstruct(X[picks]), 8, f"autoencoder rebuild from 2 numbers (error {reconstruction_error(ae, X):.2f})") 1390 pca_err = float(np.mean(np.sum((pca_reconstruction(X, 2) - X) ** 2, axis=1))) 1391 _tile(axes[2], pca_reconstruction(X, 2)[picks], 8, f"PCA rebuild from 2 numbers (error {pca_err:.2f})") 1392 fig.tight_layout() 1393 figs["reconstructions"] = fig 1394 1395 # --- 2. The autoencoder's code space: four strands ---------------------- 1396 Z = ae.encode(X) 1397 fig, ax = plt.subplots(figsize=(6, 4.4)) 1398 for k, name in enumerate(STROKE_KINDS): 1399 sel = kinds == k 1400 ax.scatter(Z[sel, 0], Z[sel, 1], s=10 + 6 * (offsets[sel] + 2.5), color=_KIND_COLOURS[k], alpha=0.8, label=name) 1401 ax.add_patch(Circle((0, 0), 2, fill=False, ls="--", color=_MUTED)) 1402 ax.annotate("where a standard normal\ndraw usually lands", (1.4, -1.4), (4, -12), color="#4b5563", 1403 arrowprops=dict(arrowstyle="->", color=_MUTED)) 1404 ax.set_xlabel("code number 1") 1405 ax.set_ylabel("code number 2") 1406 ax.set_title("A plain autoencoder's codes: one strand per stroke direction") 1407 ax.set_aspect("equal") 1408 ax.legend(frameon=False, fontsize=8, loc="upper left", bbox_to_anchor=(1.02, 1), markerscale=0.7) 1409 figs["ae_codes"] = fig 1410 1411 # --- 3. Holes: random codes inside the box decode to junk --------------- 1412 box = codes_in_box(Z, 300, seed=6) 1413 dist = distance_to_data(ae.decode(box), X) 1414 junk = dist > JUNK_DISTANCE 1415 fig, (a1, a2) = plt.subplots(1, 2, figsize=(10, 4.2), gridspec_kw={"width_ratios": [1.2, 1]}) 1416 a1.scatter(Z[:, 0], Z[:, 1], s=8, color="#111827", label="codes of real strokes") 1417 a1.scatter(box[~junk, 0], box[~junk, 1], s=14, marker="o", facecolors="none", edgecolors=_REAL, label="random code, decodes to a stroke") 1418 a1.scatter(box[junk, 0], box[junk, 1], s=14, marker="x", color=_JUNK, label="random code, decodes to junk") 1419 lo, hi = Z.min(axis=0), Z.max(axis=0) 1420 a1.add_patch(Rectangle(lo, *(hi - lo), fill=False, color=_MUTED, ls="--")) 1421 a1.set_title(f"{junk.mean():.0%} of random codes in the box land in a hole") 1422 a1.set_xlabel("code number 1") 1423 a1.set_ylabel("code number 2") 1424 a1.set_aspect("equal") 1425 a1.legend(frameon=False, fontsize=7, loc="upper center", bbox_to_anchor=(0.5, -0.15), ncol=1) 1426 order = np.argsort(-dist)[:12].tolist() + np.argsort(dist)[:12].tolist() 1427 _tile(a2, ae.decode(box[order]), 6, "decoded: the 12 worst (top) and 12 best (bottom)") 1428 fig.tight_layout() 1429 figs["holes"] = fig 1430 1431 # --- 4. The KL penalty for one code number ------------------------------- 1432 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.2)) 1433 mus = np.linspace(-3, 3, 200) 1434 a1.plot(mus, kl_to_standard_normal(mus[:, None], np.zeros((200, 1))), color=_REAL) 1435 a1.plot([0.5], [0.125], "o", color=_JUNK) 1436 a1.annotate("μ = 0.5 costs 0.125", (0.5, 0.125), (0.9, 2.2), color=_JUNK, arrowprops=dict(arrowstyle="->", color=_JUNK), zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1)) 1437 a1.set_xlabel("mean μ (spread held at σ = 1)") 1438 a1.set_ylabel("KL penalty") 1439 a1.set_title("Rent for sitting far from the centre") 1440 sigmas = np.linspace(0.05, 3, 200) 1441 a2.plot(sigmas, kl_to_standard_normal(np.zeros((200, 1)), np.log(sigmas[:, None] ** 2)), color=_REAL) 1442 a2.plot([0.5], [kl_to_standard_normal(np.zeros(1), np.log(np.array([0.25])))], "o", color=_JUNK) 1443 a2.annotate("σ = 0.5 costs 0.318", (0.5, 0.318), (1.1, 2.2), color=_JUNK, arrowprops=dict(arrowstyle="->", color=_JUNK), zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1)) 1444 a2.axvline(1, color=_MUTED, ls="--") 1445 a2.set_xlabel("spread σ (mean held at μ = 0)") 1446 a2.set_title("Rent for shrinking to a pin (or bloating)") 1447 a2.set_ylim(0, 3.5) 1448 fig.tight_layout() 1449 figs["kl_penalty"] = fig 1450 1451 # --- 5. The VAE's code space: fuzzy regions packed around 0 ------------- 1452 mu, logvar = vae.encode_distribution(X) 1453 sigma = np.exp(0.5 * logvar) 1454 fig, ax = plt.subplots(figsize=(5.4, 5)) 1455 for i in range(len(X)): 1456 ax.add_patch(Ellipse(mu[i], 2 * sigma[i, 0], 2 * sigma[i, 1], color=_KIND_COLOURS[kinds[i]], alpha=0.18, lw=0)) 1457 for k, name in enumerate(STROKE_KINDS): 1458 sel = kinds == k 1459 ax.scatter(mu[sel, 0], mu[sel, 1], s=6, color=_KIND_COLOURS[k], label=name) 1460 for r in (1, 2): 1461 ax.add_patch(Circle((0, 0), r, fill=False, ls="--", color=_MUTED)) 1462 ax.set_xlim(-3, 3) 1463 ax.set_ylim(-3, 3) 1464 ax.set_aspect("equal") 1465 ax.set_xlabel("code number 1 (μ)") 1466 ax.set_ylabel("code number 2 (μ)") 1467 ax.set_title("A VAE's codes: fuzzy regions packed inside the bell curve") 1468 ax.legend(frameon=False, fontsize=8, loc="upper left", bbox_to_anchor=(1.02, 1)) 1469 figs["vae_codes"] = fig 1470 1471 # --- 6. Samples from the standard normal: plain AE vs VAE --------------- 1472 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.2)) 1473 for ax, model, name in ((a1, ae, "plain autoencoder"), (a2, vae, "VAE")): 1474 samples = generate(model, 24, seed=5) 1475 d = np.median(distance_to_data(generate(model, 500, seed=5), X)) 1476 _tile(ax, samples, 8, f"{name}: 24 codes from N(0, I)\n(median distance to a real stroke {d:.2f})") 1477 fig.tight_layout() 1478 figs["samples"] = fig 1479 1480 # --- 7. The VAE's latent space as a map ----------------------------------- 1481 grid = np.linspace(-2.2, 2.2, 11) 1482 codes = np.array([[a, b] for b in grid[::-1] for a in grid]) 1483 fig, ax = plt.subplots(figsize=(5.4, 5.4)) 1484 _tile(ax, vae.decode(codes), len(grid), "Decoding every point of a grid from -2.2 to 2.2") 1485 ax.set_xlabel("code number 1 →") 1486 ax.set_ylabel("code number 2 →") 1487 figs["latent_grid"] = fig 1488 1489 # --- 8. A walk between two strokes --------------------------------------- 1490 a, b = _WALK 1491 fig, axes = plt.subplots(2, 1, figsize=(7, 2.6)) 1492 for ax, model, name in ((axes[0], ae, "plain autoencoder"), (axes[1], vae, "VAE")): 1493 frames = interpolate(model, X[a], X[b], steps=9) 1494 worst = distance_to_data(frames, X).max() 1495 _tile(ax, frames, 9, f"{name}: the worst stop is {worst:.2f} from any real stroke") 1496 fig.suptitle(f"Walking in a straight line between two vertical strokes (offsets {offsets[a]:.2f} and {offsets[b]:.2f})", fontsize=10) 1497 fig.tight_layout() 1498 figs["interpolation"] = fig 1499 1500 # --- 9. The β trade-off ---------------------------------------------------- 1501 rows = beta_sweep() 1502 betas = [r["beta"] for r in rows] 1503 fig = plt.figure(figsize=(10, 4.2)) 1504 a1 = fig.add_subplot(1, 2, 1) 1505 a1.plot(betas, [r["reconstruction"] for r in rows], "o-", color=_REAL, label="rebuild error") 1506 a1.plot(betas, [r["kl"] for r in rows], "o-", color="#16a34a", label="KL (information in the code)") 1507 a1.plot(betas, [r["sample_distance"] for r in rows], "o-", color=_JUNK, label="samples' distance to a real stroke") 1508 a1.set_xscale("log") 1509 a1.set_xlabel("β (weight on the KL penalty)") 1510 a1.set_title("Too little β: holes. Too much: blur, then collapse") 1511 a1.legend(frameon=False, fontsize=8) 1512 a2 = fig.add_subplot(1, 2, 2) 1513 strip = np.vstack([generate(trained_vae(beta), 8, seed=5) for beta in betas]) 1514 _tile(a2, strip, 8, "8 samples at each β (rows, top to bottom)") 1515 a2.set_yticks([(SIDE + 1) * i + SIDE / 2 - 0.5 for i in range(len(betas))], [f"β = {beta:g}" for beta in betas]) 1516 fig.tight_layout() 1517 figs["beta_tradeoff"] = fig 1518 1519 return figs 1520 1521 1522# --------------------------------------------------------------------------- 1523# 6. Narrated walkthrough 1524# --------------------------------------------------------------------------- 1525 1526 1527def demo() -> None: 1528 banner("1. A one-number code, by hand: points near the line y = 2x") 1529 ex = worked_example_line() 1530 say( 1531 """ 1532 Encoder: code = u · x with u = (1, 2)/√5, the line's direction. 1533 Decoder: rebuild = code × u. A point on the line survives the trip; 1534 a point off it comes back as its closest point on the line. 1535 """ 1536 ) 1537 table( 1538 ["point", "code", "rebuild", "squared error"], 1539 [ 1540 ("(1, 2)", ex["on_line_code"], f"({ex['on_line_rebuilt'][0]:.2f}, {ex['on_line_rebuilt'][1]:.2f})", 0.0), 1541 ("(2, 3)", ex["off_line_code"], f"({ex['off_line_rebuilt'][0]:.2f}, {ex['off_line_rebuilt'][1]:.2f})", ex["off_line_error"]), 1542 ], 1543 floatfmt=".3f", 1544 ) 1545 takeaway("An autoencoder keeps what varies most and loses the rest; training shrinks what it loses.") 1546 1547 banner("2. Squeezing 8 × 8 stroke pictures into 2 numbers") 1548 X, kinds, offsets = make_strokes() 1549 ink = float(np.mean(np.sum(X**2, axis=1))) 1550 ae = trained_autoencoder() 1551 pca_err = float(np.mean(np.sum((pca_reconstruction(X, 2) - X) ** 2, axis=1))) 1552 ae_err = reconstruction_error(ae, X) 1553 say(f"200 pictures of 64 pixels, each one stroke in one of {len(STROKE_KINDS)} directions at some offset.") 1554 table( 1555 ["2-number code", "rebuild error", "share of the picture kept"], 1556 [("PCA (flat)", pca_err, f"{1 - pca_err / ink:.0%}"), ("autoencoder (bent)", ae_err, f"{1 - ae_err / ink:.0%}")], 1557 floatfmt=".3f", 1558 ) 1559 Z = ae.encode(X) 1560 say(f"Its codes run from {Z.min(axis=0).round(1)} to {Z.max(axis=0).round(1)}: a range nothing asked for.") 1561 takeaway("Strokes lie on a curved surface in pixel space; a flat PCA sheet cannot follow it, a bent network can.") 1562 1563 banner("3. Holes: decoding random codes from a plain autoencoder") 1564 from_normal = distance_to_data(generate(ae, 500, seed=5), X) 1565 from_box = distance_to_data(ae.decode(codes_in_box(Z, 500, seed=6)), X) 1566 table( 1567 ["where the random codes come from", "median distance to a real stroke", "share that is junk"], 1568 [ 1569 ("standard normal N(0, I)", float(np.median(from_normal)), f"{np.mean(from_normal > JUNK_DISTANCE):.0%}"), 1570 ("the box around its own codes", float(np.median(from_box)), f"{np.mean(from_box > JUNK_DISTANCE):.0%}"), 1571 ], 1572 floatfmt=".2f", 1573 ) 1574 takeaway("A plain autoencoder's loss never looks between its codes, so the gaps decode to junk.") 1575 1576 banner("4. The VAE's two new pieces: the reparameterization trick and the KL rent") 1577 mu, logvar, eps = np.array([0.5]), np.log(np.array([0.25])), np.array([1.2]) 1578 say(f"μ = 0.5, σ = 0.5, ε = 1.2 -> z = μ + σ·ε = {reparameterize(mu, logvar, eps)[0]:.2f}") 1579 table( 1580 ["region", "KL rent"], 1581 [ 1582 ("μ = 0, σ = 1 (already standard)", float(kl_to_standard_normal(np.zeros(1), np.zeros(1)))), 1583 ("μ = 0.5, σ = 0.5", float(kl_to_standard_normal(mu, logvar))), 1584 ("μ = 0.5, σ = 0.05 (nearly a pin)", float(kl_to_standard_normal(mu, np.log(np.array([0.05**2]))))), 1585 ], 1586 floatfmt=".3f", 1587 ) 1588 vae = trained_vae() 1589 mu_all, logvar_all = vae.encode_distribution(X) 1590 say( 1591 f""" 1592 Trained with β = {vae.beta}, the VAE's codes centre on 1593 ({mu_all.mean(axis=0)[0]:.2f}, {mu_all.mean(axis=0)[1]:.2f}) with spread about 1594 {mu_all.std(axis=0).mean():.2f}, and each region has σ about 1595 {np.exp(0.5 * logvar_all).mean():.2f}: packed inside the bell curve. 1596 """ 1597 ) 1598 takeaway("Sampling becomes arithmetic the chain rule can pass through, and the rent packs the codes where we will sample.") 1599 1600 banner("5. Generating: draw z from N(0, I) and decode") 1601 rows = [] 1602 for model, name in ((ae, "plain autoencoder"), (vae, "VAE")): 1603 d = distance_to_data(generate(model, 500, seed=5), X) 1604 rows.append((name, float(np.median(d)), f"{np.mean(d > JUNK_DISTANCE):.0%}")) 1605 table(["model", "median distance to a real stroke", "share that is junk"], rows, floatfmt=".2f") 1606 pairs = np.random.default_rng(3).integers(0, len(X), (60, 2)) 1607 steps = { 1608 name: np.mean([largest_step(interpolate(model, X[a], X[b])) for a, b in pairs]) 1609 for model, name in ((ae, "plain autoencoder"), (vae, "VAE")) 1610 } 1611 say(f"Walking between 60 random pairs, the biggest single jump averages {steps['plain autoencoder']:.2f} (plain) vs {steps['VAE']:.2f} (VAE).") 1612 takeaway("The same random codes that give a plain autoencoder junk give a VAE mostly real-looking strokes.") 1613 1614 banner("6. The β trade-off: holes, then sharpness, then blur and collapse") 1615 table( 1616 ["β", "rebuild error", "KL", "sample distance", "sample peak ink"], 1617 [(r["beta"], r["reconstruction"], r["kl"], r["sample_distance"], r["sample_peak"]) for r in beta_sweep()], 1618 floatfmt=".3f", 1619 ) 1620 say( 1621 f""" 1622 Why blur? For a pixel that is 0 or 1 with equal chance, guessing grey 1623 costs {expected_squared_error(0.5, [0.0, 1.0]):.2f} on average and guessing 1624 either extreme costs {expected_squared_error(0.0, [0.0, 1.0]):.2f}: squared error 1625 rewards the average. At β = 3 the code carries nothing and every sample 1626 is the same grey average stroke. 1627 """ 1628 ) 1629 takeaway("β trades rebuild sharpness for a smooth, sampleable code space; latent diffusion keeps it tiny and lets diffusion generate.") 1630 1631 1632if __name__ == "__main__": 1633 demo()
1014def make_strokes(n: int = 200, seed: int = 0) -> tuple[np.ndarray, np.ndarray, np.ndarray]: 1015 """n pictures of one soft pen stroke each: (images (n, 64), kinds (n,), offsets (n,)). 1016 1017 Two hidden facts make each picture: which of the four directions the 1018 stroke runs (kinds cycle 0, 1, 2, 3) and how far it sits from the centre, 1019 a continuous offset in [-2.5, 2.5] pixels. A good 2-number code has to 1020 rediscover both from the 64 pixels alone. 1021 """ 1022 rng = np.random.default_rng(seed) 1023 kinds = np.arange(n) % len(STROKE_KINDS) 1024 offsets = rng.uniform(-2.5, 2.5, n) 1025 # Pixel centres measured from the middle of the image, so offset 0 passes through the centre. 1026 rows, cols = np.mgrid[0:SIDE, 0:SIDE] - (SIDE - 1) / 2 1027 angle = np.array(_STROKE_ANGLES)[kinds] 1028 # Signed distance from each pixel to the line: project the pixel onto the 1029 # line's perpendicular direction, then subtract how far the line is shifted. 1030 across_x, across_y = -np.sin(angle), np.cos(angle) 1031 dist = across_x[:, None, None] * cols + across_y[:, None, None] * rows - offsets[:, None, None] 1032 # Ink fades like a bell curve with distance: 1.0 on the line, about 0.25 one pixel away. 1033 images = np.exp(-(dist**2) / (2 * STROKE_WIDTH**2)) 1034 return images.reshape(n, SIDE * SIDE), kinds, offsets
n pictures of one soft pen stroke each: (images (n, 64), kinds (n,), offsets (n,)).
Two hidden facts make each picture: which of the four directions the stroke runs (kinds cycle 0, 1, 2, 3) and how far it sits from the centre, a continuous offset in [-2.5, 2.5] pixels. A good 2-number code has to rediscover both from the 64 pixels alone.
1042def worked_example_line() -> dict: 1043 """Squeeze 2-D points to one number and back, for points near y = 2x. 1044 1045 Encoder: code = u · x, with u = (1, 2)/√5, the line's direction scaled to 1046 length 1. Decoder: rebuilt = code · u. A point on the line survives the 1047 trip exactly; a point off it comes back as its closest point on the line. 1048 """ 1049 u = np.array([1.0, 2.0]) / np.sqrt(5) 1050 on, off = np.array([1.0, 2.0]), np.array([2.0, 3.0]) 1051 on_code, off_code = float(u @ on), float(u @ off) 1052 off_rebuilt = off_code * u 1053 return dict( 1054 direction=u, 1055 on_line_code=on_code, 1056 on_line_rebuilt=on_code * u, 1057 off_line_code=off_code, 1058 off_line_rebuilt=off_rebuilt, 1059 off_line_error=float(np.sum((off - off_rebuilt) ** 2)), 1060 )
Squeeze 2-D points to one number and back, for points near y = 2x.
Encoder: code = u · x, with u = (1, 2)/√5, the line's direction scaled to length 1. Decoder: rebuilt = code · u. A point on the line survives the trip exactly; a point off it comes back as its closest point on the line.
1063def pca_reconstruction(X: np.ndarray, k: int) -> np.ndarray: 1064 """Project X onto its top-k principal directions and back (see `primer.ml.embeddings.clustering.pca`). 1065 1066 PCA is the best any *flat* k-number code can do under squared error. 1067 """ 1068 mean = X.mean(axis=0) 1069 _, _, Vt = np.linalg.svd(X - mean, full_matrices=False) 1070 return (X - mean) @ Vt[:k].T @ Vt[:k] + mean
Project X onto its top-k principal directions and back (see primer.ml.embeddings.clustering.pca).
PCA is the best any flat k-number code can do under squared error.
1073def train_linear_autoencoder(X: np.ndarray, n_code: int = 1, steps: int = 2000, lr: float = 0.02, seed: int = 0): 1074 """An autoencoder with no bends: code = x W_enc, rebuilt = code W_dec. Returns (W_enc, W_dec). 1075 1076 Trained by plain gradient descent on the squared error of centred data. 1077 With nothing nonlinear in it, the best it can do is PCA's answer: it ends 1078 up spanning the same directions as the top principal components. 1079 """ 1080 rng = np.random.default_rng(seed) 1081 Xc = X - X.mean(axis=0) 1082 n, d = Xc.shape 1083 W_enc = rng.normal(0, 0.1, (d, n_code)) # (d, k) 1084 W_dec = rng.normal(0, 0.1, (n_code, d)) # (k, d) 1085 for _ in range(steps): 1086 Z = Xc @ W_enc # (n, k) codes 1087 R = Z @ W_dec - Xc # (n, d) what the rebuild got wrong 1088 # Loss = mean over points of the summed squared error; these are its exact gradients. 1089 g_dec = 2 / n * Z.T @ R 1090 g_enc = 2 / n * Xc.T @ (R @ W_dec.T) 1091 W_enc -= lr * g_enc 1092 W_dec -= lr * g_dec 1093 return W_enc, W_dec
An autoencoder with no bends: code = x W_enc, rebuilt = code W_dec. Returns (W_enc, W_dec).
Trained by plain gradient descent on the squared error of centred data. With nothing nonlinear in it, the best it can do is PCA's answer: it ends up spanning the same directions as the top principal components.
1106class Autoencoder: 1107 """64 pixels -> 32 hidden -> a code of `n_code` numbers -> 32 hidden -> 64 pixels. 1108 1109 The encoder and the decoder are each a two-layer network like the one in 1110 `primer.ml.neural_net`. Training asks only one thing: rebuild the input. 1111 """ 1112 1113 def __init__(self, n_in: int = SIDE * SIDE, n_hidden: int = 32, n_code: int = 2, seed: int = 0): 1114 self.n_in, self.n_hidden, self.n_code = n_in, n_hidden, n_code 1115 rng = np.random.default_rng(seed) 1116 head = self._head_width() 1117 # Variance 1/fan_in keeps tanh out of its flat regions at the start (see primer.ml.deep_nets). 1118 self.params = { 1119 "W1": rng.normal(0, 1 / np.sqrt(n_in), (n_in, n_hidden)), # encoder 1120 "b1": np.zeros(n_hidden), 1121 "W2": rng.normal(0, 1 / np.sqrt(n_hidden), (n_hidden, head)), 1122 "b2": np.zeros(head), 1123 "W3": rng.normal(0, 1 / np.sqrt(n_code), (n_code, n_hidden)), # decoder 1124 "b3": np.zeros(n_hidden), 1125 "W4": rng.normal(0, 1 / np.sqrt(n_hidden), (n_hidden, n_in)), 1126 "b4": np.zeros(n_in), 1127 } 1128 1129 def _head_width(self) -> int: 1130 return self.n_code # the encoder's last layer outputs the code itself 1131 1132 def _encoder(self, X: np.ndarray) -> tuple[np.ndarray, np.ndarray]: 1133 p = self.params 1134 h = np.tanh(X @ p["W1"] + p["b1"]) # (n, hidden) 1135 return h, h @ p["W2"] + p["b2"] # (n, head): no squashing, a code may be any number 1136 1137 def _decoder(self, Z: np.ndarray) -> tuple[np.ndarray, np.ndarray]: 1138 p = self.params 1139 g = np.tanh(Z @ p["W3"] + p["b3"]) # (n, hidden) 1140 return g, sigmoid(g @ p["W4"] + p["b4"]) # (n, 64) pixels in (0, 1) 1141 1142 def encode(self, X: np.ndarray) -> np.ndarray: 1143 """Images (n, 64) -> codes (n, n_code).""" 1144 return self._encoder(X)[1] 1145 1146 def decode(self, Z: np.ndarray) -> np.ndarray: 1147 """Codes (n, n_code) -> images (n, 64). Works on any code, seen in training or not.""" 1148 return self._decoder(Z)[1] 1149 1150 def reconstruct(self, X: np.ndarray) -> np.ndarray: 1151 return self.decode(self.encode(X)) 1152 1153 def _backward_decoder(self, Z, g, X_hat, X) -> tuple[dict, np.ndarray]: 1154 """Gradients of the mean summed squared error for the decoder, and the error signal at the code.""" 1155 p, n = self.params, len(X) 1156 d_out = 2 * (X_hat - X) / n * X_hat * (1 - X_hat) # through the squared error, then the sigmoid 1157 d_g = d_out @ p["W4"].T * (1 - g**2) # back through W4, then tanh' = 1 - tanh² 1158 grads = {"W4": g.T @ d_out, "b4": d_out.sum(axis=0), "W3": Z.T @ d_g, "b3": d_g.sum(axis=0)} 1159 return grads, d_g @ p["W3"].T # (n, n_code): how the loss changes with each code number 1160 1161 def _backward_encoder(self, X, h, d_head) -> dict: 1162 p = self.params 1163 d_h = d_head @ p["W2"].T * (1 - h**2) 1164 return {"W2": h.T @ d_head, "b2": d_head.sum(axis=0), "W1": X.T @ d_h, "b1": d_h.sum(axis=0)} 1165 1166 def loss_and_gradients(self, X: np.ndarray, noise: np.ndarray | None = None) -> tuple[dict, dict]: 1167 """Reconstruction loss (squared error summed over pixels, averaged over images) and every gradient. 1168 1169 `noise` is ignored; it is accepted so `train` can drive this and `VAE` the same way. 1170 """ 1171 h, Z = self._encoder(X) 1172 g, X_hat = self._decoder(Z) 1173 rec = float(np.sum((X_hat - X) ** 2) / len(X)) 1174 grads, d_Z = self._backward_decoder(Z, g, X_hat, X) 1175 grads |= self._backward_encoder(X, h, d_Z) 1176 return {"loss": rec, "reconstruction": rec, "kl": 0.0}, grads
64 pixels -> 32 hidden -> a code of n_code numbers -> 32 hidden -> 64 pixels.
The encoder and the decoder are each a two-layer network like the one in
primer.ml.neural_net. Training asks only one thing: rebuild the input.
1113 def __init__(self, n_in: int = SIDE * SIDE, n_hidden: int = 32, n_code: int = 2, seed: int = 0): 1114 self.n_in, self.n_hidden, self.n_code = n_in, n_hidden, n_code 1115 rng = np.random.default_rng(seed) 1116 head = self._head_width() 1117 # Variance 1/fan_in keeps tanh out of its flat regions at the start (see primer.ml.deep_nets). 1118 self.params = { 1119 "W1": rng.normal(0, 1 / np.sqrt(n_in), (n_in, n_hidden)), # encoder 1120 "b1": np.zeros(n_hidden), 1121 "W2": rng.normal(0, 1 / np.sqrt(n_hidden), (n_hidden, head)), 1122 "b2": np.zeros(head), 1123 "W3": rng.normal(0, 1 / np.sqrt(n_code), (n_code, n_hidden)), # decoder 1124 "b3": np.zeros(n_hidden), 1125 "W4": rng.normal(0, 1 / np.sqrt(n_hidden), (n_hidden, n_in)), 1126 "b4": np.zeros(n_in), 1127 }
1142 def encode(self, X: np.ndarray) -> np.ndarray: 1143 """Images (n, 64) -> codes (n, n_code).""" 1144 return self._encoder(X)[1]
Images (n, 64) -> codes (n, n_code).
1146 def decode(self, Z: np.ndarray) -> np.ndarray: 1147 """Codes (n, n_code) -> images (n, 64). Works on any code, seen in training or not.""" 1148 return self._decoder(Z)[1]
Codes (n, n_code) -> images (n, 64). Works on any code, seen in training or not.
1166 def loss_and_gradients(self, X: np.ndarray, noise: np.ndarray | None = None) -> tuple[dict, dict]: 1167 """Reconstruction loss (squared error summed over pixels, averaged over images) and every gradient. 1168 1169 `noise` is ignored; it is accepted so `train` can drive this and `VAE` the same way. 1170 """ 1171 h, Z = self._encoder(X) 1172 g, X_hat = self._decoder(Z) 1173 rec = float(np.sum((X_hat - X) ** 2) / len(X)) 1174 grads, d_Z = self._backward_decoder(Z, g, X_hat, X) 1175 grads |= self._backward_encoder(X, h, d_Z) 1176 return {"loss": rec, "reconstruction": rec, "kl": 0.0}, grads
1179class VAE(Autoencoder): 1180 """A variational autoencoder: the encoder outputs a fuzzy region, not a point. 1181 1182 For each image the encoder gives a mean μ and a log-variance log σ² per 1183 code number. Training draws a code from that region with the 1184 reparameterization trick and adds `beta` times the KL penalty that pulls 1185 every region towards the standard normal distribution. 1186 """ 1187 1188 def __init__(self, n_in: int = SIDE * SIDE, n_hidden: int = 32, n_code: int = 2, beta: float = 0.3, seed: int = 0): 1189 self.beta = beta 1190 super().__init__(n_in, n_hidden, n_code, seed) 1191 1192 def _head_width(self) -> int: 1193 return 2 * self.n_code # μ and log σ² for every code number 1194 1195 def encode_distribution(self, X: np.ndarray) -> tuple[np.ndarray, np.ndarray]: 1196 """Images -> (μ, log σ²), each (n, n_code).""" 1197 head = self._encoder(X)[1] 1198 return head[:, : self.n_code], head[:, self.n_code :] 1199 1200 def encode(self, X: np.ndarray) -> np.ndarray: 1201 """The centre μ of each image's region: the single best code to decode it from.""" 1202 return self.encode_distribution(X)[0] 1203 1204 def loss_and_gradients(self, X: np.ndarray, noise: np.ndarray | None = None) -> tuple[dict, dict]: 1205 """Reconstruction + β·KL, with gradients through the sampling step. 1206 1207 `noise` is ε, shape (n, n_code), drawn from the standard normal by the 1208 caller. Passing it in keeps the step deterministic, which is what lets 1209 the tests check every gradient against finite differences. 1210 """ 1211 n = len(X) 1212 noise = np.zeros((n, self.n_code)) if noise is None else noise 1213 h, head = self._encoder(X) 1214 mu, logvar = head[:, : self.n_code], head[:, self.n_code :] 1215 sigma = np.exp(0.5 * logvar) 1216 Z = reparameterize(mu, logvar, noise) # μ + σ·ε: randomness enters only through ε 1217 g, X_hat = self._decoder(Z) 1218 rec = float(np.sum((X_hat - X) ** 2) / n) 1219 kl = float(np.sum(kl_to_standard_normal(mu, logvar)) / n) 1220 grads, d_Z = self._backward_decoder(Z, g, X_hat, X) 1221 # Through z = μ + σ·ε: ∂z/∂μ = 1 and ∂z/∂σ = ε, with ∂σ/∂(log σ²) = σ/2. 1222 # The KL term adds its own pull: ∂KL/∂μ = μ and ∂KL/∂(log σ²) = (σ² - 1)/2. 1223 d_mu = d_Z + self.beta * mu / n 1224 d_logvar = d_Z * noise * sigma / 2 + self.beta * (sigma**2 - 1) / (2 * n) 1225 grads |= self._backward_encoder(X, h, np.hstack([d_mu, d_logvar])) 1226 return {"loss": rec + self.beta * kl, "reconstruction": rec, "kl": kl}, grads
A variational autoencoder: the encoder outputs a fuzzy region, not a point.
For each image the encoder gives a mean μ and a log-variance log σ² per
code number. Training draws a code from that region with the
reparameterization trick and adds beta times the KL penalty that pulls
every region towards the standard normal distribution.
1195 def encode_distribution(self, X: np.ndarray) -> tuple[np.ndarray, np.ndarray]: 1196 """Images -> (μ, log σ²), each (n, n_code).""" 1197 head = self._encoder(X)[1] 1198 return head[:, : self.n_code], head[:, self.n_code :]
Images -> (μ, log σ²), each (n, n_code).
1200 def encode(self, X: np.ndarray) -> np.ndarray: 1201 """The centre μ of each image's region: the single best code to decode it from.""" 1202 return self.encode_distribution(X)[0]
The centre μ of each image's region: the single best code to decode it from.
1204 def loss_and_gradients(self, X: np.ndarray, noise: np.ndarray | None = None) -> tuple[dict, dict]: 1205 """Reconstruction + β·KL, with gradients through the sampling step. 1206 1207 `noise` is ε, shape (n, n_code), drawn from the standard normal by the 1208 caller. Passing it in keeps the step deterministic, which is what lets 1209 the tests check every gradient against finite differences. 1210 """ 1211 n = len(X) 1212 noise = np.zeros((n, self.n_code)) if noise is None else noise 1213 h, head = self._encoder(X) 1214 mu, logvar = head[:, : self.n_code], head[:, self.n_code :] 1215 sigma = np.exp(0.5 * logvar) 1216 Z = reparameterize(mu, logvar, noise) # μ + σ·ε: randomness enters only through ε 1217 g, X_hat = self._decoder(Z) 1218 rec = float(np.sum((X_hat - X) ** 2) / n) 1219 kl = float(np.sum(kl_to_standard_normal(mu, logvar)) / n) 1220 grads, d_Z = self._backward_decoder(Z, g, X_hat, X) 1221 # Through z = μ + σ·ε: ∂z/∂μ = 1 and ∂z/∂σ = ε, with ∂σ/∂(log σ²) = σ/2. 1222 # The KL term adds its own pull: ∂KL/∂μ = μ and ∂KL/∂(log σ²) = (σ² - 1)/2. 1223 d_mu = d_Z + self.beta * mu / n 1224 d_logvar = d_Z * noise * sigma / 2 + self.beta * (sigma**2 - 1) / (2 * n) 1225 grads |= self._backward_encoder(X, h, np.hstack([d_mu, d_logvar])) 1226 return {"loss": rec + self.beta * kl, "reconstruction": rec, "kl": kl}, grads
Reconstruction + β·KL, with gradients through the sampling step.
noise is ε, shape (n, n_code), drawn from the standard normal by the
caller. Passing it in keeps the step deterministic, which is what lets
the tests check every gradient against finite differences.
Inherited Members
1229def reparameterize(mu: np.ndarray, logvar: np.ndarray, eps: np.ndarray) -> np.ndarray: 1230 """z = μ + σ·ε with σ = exp(½ log σ²): a draw from N(μ, σ²) written as a plain sum. 1231 1232 The encoder predicts log σ² rather than σ because a log can be any number, 1233 positive or negative, while σ must stay positive; exp takes care of that. 1234 """ 1235 return mu + np.exp(0.5 * logvar) * eps
z = μ + σ·ε with σ = exp(½ log σ²): a draw from N(μ, σ²) written as a plain sum.
The encoder predicts log σ² rather than σ because a log can be any number, positive or negative, while σ must stay positive; exp takes care of that.
1238def kl_to_standard_normal(mu: np.ndarray, logvar: np.ndarray) -> np.ndarray: 1239 """KL(N(μ, σ²) ‖ N(0, 1)) summed over the last axis: ½ Σ (μ² + σ² - log σ² - 1).""" 1240 return 0.5 * np.sum(mu**2 + np.exp(logvar) - logvar - 1, axis=-1)
KL(N(μ, σ²) ‖ N(0, 1)) summed over the last axis: ½ Σ (μ² + σ² - log σ² - 1).
1243def train(model: Autoencoder, X: np.ndarray, steps: int = 1500, lr: float = 0.01, seed: int = 0) -> list[dict]: 1244 """Full-batch training with Adam (see `primer.ml.optimizers`). Returns the losses at every step.""" 1245 rng = np.random.default_rng(seed) 1246 optimizers = {name: Adam(lr=lr) for name in model.params} 1247 history = [] 1248 for _ in range(steps): 1249 # Fresh ε every step, so each image's code lands somewhere new in its region. 1250 noise = rng.standard_normal((len(X), model.n_code)) 1251 losses, grads = model.loss_and_gradients(X, noise) 1252 for name in model.params: 1253 model.params[name] = optimizers[name].step(model.params[name], grads[name]) 1254 history.append(losses) 1255 return history
Full-batch training with Adam (see primer.ml.optimizers). Returns the losses at every step.
1258@lru_cache(maxsize=None) 1259def trained_autoencoder(steps: int = 1500) -> Autoencoder: 1260 """The lesson's plain autoencoder, trained once on `make_strokes` and reused.""" 1261 model = Autoencoder() 1262 train(model, make_strokes()[0], steps=steps) 1263 return model
The lesson's plain autoencoder, trained once on make_strokes and reused.
1266@lru_cache(maxsize=None) 1267def trained_vae(beta: float = 0.3, steps: int = 1500) -> VAE: 1268 """The lesson's VAE at a given β, trained once on `make_strokes` and reused.""" 1269 model = VAE(beta=beta) 1270 train(model, make_strokes()[0], steps=steps) 1271 return model
The lesson's VAE at a given β, trained once on make_strokes and reused.
1279def reconstruction_error(model: Autoencoder, X: np.ndarray) -> float: 1280 """Squared error summed over the 64 pixels, averaged over images.""" 1281 return float(np.mean(np.sum((model.reconstruct(X) - X) ** 2, axis=1)))
Squared error summed over the 64 pixels, averaged over images.
1284def distance_to_data(images: np.ndarray, X: np.ndarray) -> np.ndarray: 1285 """For each image, the squared distance to the nearest training image. Large means junk.""" 1286 # ‖a - b‖² = ‖a‖² + ‖b‖² - 2 a·b, for every pair at once: (n_images, n_train). 1287 d = np.sum(images**2, axis=1)[:, None] + np.sum(X**2, axis=1)[None, :] - 2 * images @ X.T 1288 return np.maximum(d.min(axis=1), 0.0)
For each image, the squared distance to the nearest training image. Large means junk.
1291def generate(model: Autoencoder, n: int, seed: int = 0) -> np.ndarray: 1292 """Draw n codes from the standard normal and decode them: brand-new images.""" 1293 Z = np.random.default_rng(seed).standard_normal((n, model.n_code)) 1294 return model.decode(Z)
Draw n codes from the standard normal and decode them: brand-new images.
1297def codes_in_box(codes: np.ndarray, n: int, seed: int = 0) -> np.ndarray: 1298 """n codes drawn evenly from the smallest box that holds every given code.""" 1299 rng = np.random.default_rng(seed) 1300 return rng.uniform(codes.min(axis=0), codes.max(axis=0), (n, codes.shape[1]))
n codes drawn evenly from the smallest box that holds every given code.
1303def interpolate(model: Autoencoder, x_a: np.ndarray, x_b: np.ndarray, steps: int = 9) -> np.ndarray: 1304 """Encode two images, walk a straight line between their codes, decode every stop: (steps, 64).""" 1305 z_a, z_b = model.encode(np.stack([x_a, x_b])) 1306 t = np.linspace(0, 1, steps)[:, None] # 0 at image a, 1 at image b 1307 return model.decode((1 - t) * z_a + t * z_b)
Encode two images, walk a straight line between their codes, decode every stop: (steps, 64).
1310def largest_step(frames: np.ndarray) -> float: 1311 """The biggest change (Euclidean distance in pixels) between consecutive frames of a walk.""" 1312 return float(np.max(np.linalg.norm(np.diff(frames, axis=0), axis=1)))
The biggest change (Euclidean distance in pixels) between consecutive frames of a walk.
1315def expected_squared_error(guess: float, outcomes) -> float: 1316 """Average squared error of one guess against equally likely outcomes.""" 1317 return float(np.mean([(o - guess) ** 2 for o in outcomes]))
Average squared error of one guess against equally likely outcomes.
1320def beta_sweep(betas=(0.01, 0.03, 0.1, 0.3, 1.0, 3.0)) -> list[dict]: 1321 """Train a VAE at each β and measure rebuild quality, KL, and how real its samples look.""" 1322 X = make_strokes()[0] 1323 rows = [] 1324 for beta in betas: 1325 model = trained_vae(beta) 1326 samples = generate(model, 500, seed=5) 1327 rows.append( 1328 dict( 1329 beta=beta, 1330 reconstruction=reconstruction_error(model, X), 1331 kl=float(np.mean(kl_to_standard_normal(*model.encode_distribution(X)))), 1332 sample_distance=float(np.median(distance_to_data(samples, X))), 1333 sample_peak=float(samples.max(axis=1).mean()), 1334 ) 1335 ) 1336 return rows
Train a VAE at each β and measure rebuild quality, KL, and how real its samples look.
1373def figures() -> dict: 1374 """Plot this lesson's data. matplotlib is imported here, and only here, 1375 so the lesson itself needs nothing beyond NumPy.""" 1376 import matplotlib 1377 1378 matplotlib.use("Agg") 1379 import matplotlib.pyplot as plt 1380 from matplotlib.patches import Circle, Ellipse, Rectangle 1381 1382 X, kinds, offsets = make_strokes() 1383 ae, vae = trained_autoencoder(), trained_vae() 1384 figs = {} 1385 1386 # --- 1. Rebuilds: autoencoder vs PCA, both with 2 numbers ---------------- 1387 picks = _pick_examples(kinds, offsets, per_kind=2) 1388 fig, axes = plt.subplots(3, 1, figsize=(7, 3.6)) 1389 _tile(axes[0], X[picks], 8, "the original pictures (64 pixels each)") 1390 _tile(axes[1], ae.reconstruct(X[picks]), 8, f"autoencoder rebuild from 2 numbers (error {reconstruction_error(ae, X):.2f})") 1391 pca_err = float(np.mean(np.sum((pca_reconstruction(X, 2) - X) ** 2, axis=1))) 1392 _tile(axes[2], pca_reconstruction(X, 2)[picks], 8, f"PCA rebuild from 2 numbers (error {pca_err:.2f})") 1393 fig.tight_layout() 1394 figs["reconstructions"] = fig 1395 1396 # --- 2. The autoencoder's code space: four strands ---------------------- 1397 Z = ae.encode(X) 1398 fig, ax = plt.subplots(figsize=(6, 4.4)) 1399 for k, name in enumerate(STROKE_KINDS): 1400 sel = kinds == k 1401 ax.scatter(Z[sel, 0], Z[sel, 1], s=10 + 6 * (offsets[sel] + 2.5), color=_KIND_COLOURS[k], alpha=0.8, label=name) 1402 ax.add_patch(Circle((0, 0), 2, fill=False, ls="--", color=_MUTED)) 1403 ax.annotate("where a standard normal\ndraw usually lands", (1.4, -1.4), (4, -12), color="#4b5563", 1404 arrowprops=dict(arrowstyle="->", color=_MUTED)) 1405 ax.set_xlabel("code number 1") 1406 ax.set_ylabel("code number 2") 1407 ax.set_title("A plain autoencoder's codes: one strand per stroke direction") 1408 ax.set_aspect("equal") 1409 ax.legend(frameon=False, fontsize=8, loc="upper left", bbox_to_anchor=(1.02, 1), markerscale=0.7) 1410 figs["ae_codes"] = fig 1411 1412 # --- 3. Holes: random codes inside the box decode to junk --------------- 1413 box = codes_in_box(Z, 300, seed=6) 1414 dist = distance_to_data(ae.decode(box), X) 1415 junk = dist > JUNK_DISTANCE 1416 fig, (a1, a2) = plt.subplots(1, 2, figsize=(10, 4.2), gridspec_kw={"width_ratios": [1.2, 1]}) 1417 a1.scatter(Z[:, 0], Z[:, 1], s=8, color="#111827", label="codes of real strokes") 1418 a1.scatter(box[~junk, 0], box[~junk, 1], s=14, marker="o", facecolors="none", edgecolors=_REAL, label="random code, decodes to a stroke") 1419 a1.scatter(box[junk, 0], box[junk, 1], s=14, marker="x", color=_JUNK, label="random code, decodes to junk") 1420 lo, hi = Z.min(axis=0), Z.max(axis=0) 1421 a1.add_patch(Rectangle(lo, *(hi - lo), fill=False, color=_MUTED, ls="--")) 1422 a1.set_title(f"{junk.mean():.0%} of random codes in the box land in a hole") 1423 a1.set_xlabel("code number 1") 1424 a1.set_ylabel("code number 2") 1425 a1.set_aspect("equal") 1426 a1.legend(frameon=False, fontsize=7, loc="upper center", bbox_to_anchor=(0.5, -0.15), ncol=1) 1427 order = np.argsort(-dist)[:12].tolist() + np.argsort(dist)[:12].tolist() 1428 _tile(a2, ae.decode(box[order]), 6, "decoded: the 12 worst (top) and 12 best (bottom)") 1429 fig.tight_layout() 1430 figs["holes"] = fig 1431 1432 # --- 4. The KL penalty for one code number ------------------------------- 1433 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.2)) 1434 mus = np.linspace(-3, 3, 200) 1435 a1.plot(mus, kl_to_standard_normal(mus[:, None], np.zeros((200, 1))), color=_REAL) 1436 a1.plot([0.5], [0.125], "o", color=_JUNK) 1437 a1.annotate("μ = 0.5 costs 0.125", (0.5, 0.125), (0.9, 2.2), color=_JUNK, arrowprops=dict(arrowstyle="->", color=_JUNK), zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1)) 1438 a1.set_xlabel("mean μ (spread held at σ = 1)") 1439 a1.set_ylabel("KL penalty") 1440 a1.set_title("Rent for sitting far from the centre") 1441 sigmas = np.linspace(0.05, 3, 200) 1442 a2.plot(sigmas, kl_to_standard_normal(np.zeros((200, 1)), np.log(sigmas[:, None] ** 2)), color=_REAL) 1443 a2.plot([0.5], [kl_to_standard_normal(np.zeros(1), np.log(np.array([0.25])))], "o", color=_JUNK) 1444 a2.annotate("σ = 0.5 costs 0.318", (0.5, 0.318), (1.1, 2.2), color=_JUNK, arrowprops=dict(arrowstyle="->", color=_JUNK), zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1)) 1445 a2.axvline(1, color=_MUTED, ls="--") 1446 a2.set_xlabel("spread σ (mean held at μ = 0)") 1447 a2.set_title("Rent for shrinking to a pin (or bloating)") 1448 a2.set_ylim(0, 3.5) 1449 fig.tight_layout() 1450 figs["kl_penalty"] = fig 1451 1452 # --- 5. The VAE's code space: fuzzy regions packed around 0 ------------- 1453 mu, logvar = vae.encode_distribution(X) 1454 sigma = np.exp(0.5 * logvar) 1455 fig, ax = plt.subplots(figsize=(5.4, 5)) 1456 for i in range(len(X)): 1457 ax.add_patch(Ellipse(mu[i], 2 * sigma[i, 0], 2 * sigma[i, 1], color=_KIND_COLOURS[kinds[i]], alpha=0.18, lw=0)) 1458 for k, name in enumerate(STROKE_KINDS): 1459 sel = kinds == k 1460 ax.scatter(mu[sel, 0], mu[sel, 1], s=6, color=_KIND_COLOURS[k], label=name) 1461 for r in (1, 2): 1462 ax.add_patch(Circle((0, 0), r, fill=False, ls="--", color=_MUTED)) 1463 ax.set_xlim(-3, 3) 1464 ax.set_ylim(-3, 3) 1465 ax.set_aspect("equal") 1466 ax.set_xlabel("code number 1 (μ)") 1467 ax.set_ylabel("code number 2 (μ)") 1468 ax.set_title("A VAE's codes: fuzzy regions packed inside the bell curve") 1469 ax.legend(frameon=False, fontsize=8, loc="upper left", bbox_to_anchor=(1.02, 1)) 1470 figs["vae_codes"] = fig 1471 1472 # --- 6. Samples from the standard normal: plain AE vs VAE --------------- 1473 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.2)) 1474 for ax, model, name in ((a1, ae, "plain autoencoder"), (a2, vae, "VAE")): 1475 samples = generate(model, 24, seed=5) 1476 d = np.median(distance_to_data(generate(model, 500, seed=5), X)) 1477 _tile(ax, samples, 8, f"{name}: 24 codes from N(0, I)\n(median distance to a real stroke {d:.2f})") 1478 fig.tight_layout() 1479 figs["samples"] = fig 1480 1481 # --- 7. The VAE's latent space as a map ----------------------------------- 1482 grid = np.linspace(-2.2, 2.2, 11) 1483 codes = np.array([[a, b] for b in grid[::-1] for a in grid]) 1484 fig, ax = plt.subplots(figsize=(5.4, 5.4)) 1485 _tile(ax, vae.decode(codes), len(grid), "Decoding every point of a grid from -2.2 to 2.2") 1486 ax.set_xlabel("code number 1 →") 1487 ax.set_ylabel("code number 2 →") 1488 figs["latent_grid"] = fig 1489 1490 # --- 8. A walk between two strokes --------------------------------------- 1491 a, b = _WALK 1492 fig, axes = plt.subplots(2, 1, figsize=(7, 2.6)) 1493 for ax, model, name in ((axes[0], ae, "plain autoencoder"), (axes[1], vae, "VAE")): 1494 frames = interpolate(model, X[a], X[b], steps=9) 1495 worst = distance_to_data(frames, X).max() 1496 _tile(ax, frames, 9, f"{name}: the worst stop is {worst:.2f} from any real stroke") 1497 fig.suptitle(f"Walking in a straight line between two vertical strokes (offsets {offsets[a]:.2f} and {offsets[b]:.2f})", fontsize=10) 1498 fig.tight_layout() 1499 figs["interpolation"] = fig 1500 1501 # --- 9. The β trade-off ---------------------------------------------------- 1502 rows = beta_sweep() 1503 betas = [r["beta"] for r in rows] 1504 fig = plt.figure(figsize=(10, 4.2)) 1505 a1 = fig.add_subplot(1, 2, 1) 1506 a1.plot(betas, [r["reconstruction"] for r in rows], "o-", color=_REAL, label="rebuild error") 1507 a1.plot(betas, [r["kl"] for r in rows], "o-", color="#16a34a", label="KL (information in the code)") 1508 a1.plot(betas, [r["sample_distance"] for r in rows], "o-", color=_JUNK, label="samples' distance to a real stroke") 1509 a1.set_xscale("log") 1510 a1.set_xlabel("β (weight on the KL penalty)") 1511 a1.set_title("Too little β: holes. Too much: blur, then collapse") 1512 a1.legend(frameon=False, fontsize=8) 1513 a2 = fig.add_subplot(1, 2, 2) 1514 strip = np.vstack([generate(trained_vae(beta), 8, seed=5) for beta in betas]) 1515 _tile(a2, strip, 8, "8 samples at each β (rows, top to bottom)") 1516 a2.set_yticks([(SIDE + 1) * i + SIDE / 2 - 0.5 for i in range(len(betas))], [f"β = {beta:g}" for beta in betas]) 1517 fig.tight_layout() 1518 figs["beta_tradeoff"] = fig 1519 1520 return figs
Plot this lesson's data. matplotlib is imported here, and only here, so the lesson itself needs nothing beyond NumPy.
1528def demo() -> None: 1529 banner("1. A one-number code, by hand: points near the line y = 2x") 1530 ex = worked_example_line() 1531 say( 1532 """ 1533 Encoder: code = u · x with u = (1, 2)/√5, the line's direction. 1534 Decoder: rebuild = code × u. A point on the line survives the trip; 1535 a point off it comes back as its closest point on the line. 1536 """ 1537 ) 1538 table( 1539 ["point", "code", "rebuild", "squared error"], 1540 [ 1541 ("(1, 2)", ex["on_line_code"], f"({ex['on_line_rebuilt'][0]:.2f}, {ex['on_line_rebuilt'][1]:.2f})", 0.0), 1542 ("(2, 3)", ex["off_line_code"], f"({ex['off_line_rebuilt'][0]:.2f}, {ex['off_line_rebuilt'][1]:.2f})", ex["off_line_error"]), 1543 ], 1544 floatfmt=".3f", 1545 ) 1546 takeaway("An autoencoder keeps what varies most and loses the rest; training shrinks what it loses.") 1547 1548 banner("2. Squeezing 8 × 8 stroke pictures into 2 numbers") 1549 X, kinds, offsets = make_strokes() 1550 ink = float(np.mean(np.sum(X**2, axis=1))) 1551 ae = trained_autoencoder() 1552 pca_err = float(np.mean(np.sum((pca_reconstruction(X, 2) - X) ** 2, axis=1))) 1553 ae_err = reconstruction_error(ae, X) 1554 say(f"200 pictures of 64 pixels, each one stroke in one of {len(STROKE_KINDS)} directions at some offset.") 1555 table( 1556 ["2-number code", "rebuild error", "share of the picture kept"], 1557 [("PCA (flat)", pca_err, f"{1 - pca_err / ink:.0%}"), ("autoencoder (bent)", ae_err, f"{1 - ae_err / ink:.0%}")], 1558 floatfmt=".3f", 1559 ) 1560 Z = ae.encode(X) 1561 say(f"Its codes run from {Z.min(axis=0).round(1)} to {Z.max(axis=0).round(1)}: a range nothing asked for.") 1562 takeaway("Strokes lie on a curved surface in pixel space; a flat PCA sheet cannot follow it, a bent network can.") 1563 1564 banner("3. Holes: decoding random codes from a plain autoencoder") 1565 from_normal = distance_to_data(generate(ae, 500, seed=5), X) 1566 from_box = distance_to_data(ae.decode(codes_in_box(Z, 500, seed=6)), X) 1567 table( 1568 ["where the random codes come from", "median distance to a real stroke", "share that is junk"], 1569 [ 1570 ("standard normal N(0, I)", float(np.median(from_normal)), f"{np.mean(from_normal > JUNK_DISTANCE):.0%}"), 1571 ("the box around its own codes", float(np.median(from_box)), f"{np.mean(from_box > JUNK_DISTANCE):.0%}"), 1572 ], 1573 floatfmt=".2f", 1574 ) 1575 takeaway("A plain autoencoder's loss never looks between its codes, so the gaps decode to junk.") 1576 1577 banner("4. The VAE's two new pieces: the reparameterization trick and the KL rent") 1578 mu, logvar, eps = np.array([0.5]), np.log(np.array([0.25])), np.array([1.2]) 1579 say(f"μ = 0.5, σ = 0.5, ε = 1.2 -> z = μ + σ·ε = {reparameterize(mu, logvar, eps)[0]:.2f}") 1580 table( 1581 ["region", "KL rent"], 1582 [ 1583 ("μ = 0, σ = 1 (already standard)", float(kl_to_standard_normal(np.zeros(1), np.zeros(1)))), 1584 ("μ = 0.5, σ = 0.5", float(kl_to_standard_normal(mu, logvar))), 1585 ("μ = 0.5, σ = 0.05 (nearly a pin)", float(kl_to_standard_normal(mu, np.log(np.array([0.05**2]))))), 1586 ], 1587 floatfmt=".3f", 1588 ) 1589 vae = trained_vae() 1590 mu_all, logvar_all = vae.encode_distribution(X) 1591 say( 1592 f""" 1593 Trained with β = {vae.beta}, the VAE's codes centre on 1594 ({mu_all.mean(axis=0)[0]:.2f}, {mu_all.mean(axis=0)[1]:.2f}) with spread about 1595 {mu_all.std(axis=0).mean():.2f}, and each region has σ about 1596 {np.exp(0.5 * logvar_all).mean():.2f}: packed inside the bell curve. 1597 """ 1598 ) 1599 takeaway("Sampling becomes arithmetic the chain rule can pass through, and the rent packs the codes where we will sample.") 1600 1601 banner("5. Generating: draw z from N(0, I) and decode") 1602 rows = [] 1603 for model, name in ((ae, "plain autoencoder"), (vae, "VAE")): 1604 d = distance_to_data(generate(model, 500, seed=5), X) 1605 rows.append((name, float(np.median(d)), f"{np.mean(d > JUNK_DISTANCE):.0%}")) 1606 table(["model", "median distance to a real stroke", "share that is junk"], rows, floatfmt=".2f") 1607 pairs = np.random.default_rng(3).integers(0, len(X), (60, 2)) 1608 steps = { 1609 name: np.mean([largest_step(interpolate(model, X[a], X[b])) for a, b in pairs]) 1610 for model, name in ((ae, "plain autoencoder"), (vae, "VAE")) 1611 } 1612 say(f"Walking between 60 random pairs, the biggest single jump averages {steps['plain autoencoder']:.2f} (plain) vs {steps['VAE']:.2f} (VAE).") 1613 takeaway("The same random codes that give a plain autoencoder junk give a VAE mostly real-looking strokes.") 1614 1615 banner("6. The β trade-off: holes, then sharpness, then blur and collapse") 1616 table( 1617 ["β", "rebuild error", "KL", "sample distance", "sample peak ink"], 1618 [(r["beta"], r["reconstruction"], r["kl"], r["sample_distance"], r["sample_peak"]) for r in beta_sweep()], 1619 floatfmt=".3f", 1620 ) 1621 say( 1622 f""" 1623 Why blur? For a pixel that is 0 or 1 with equal chance, guessing grey 1624 costs {expected_squared_error(0.5, [0.0, 1.0]):.2f} on average and guessing 1625 either extreme costs {expected_squared_error(0.0, [0.0, 1.0]):.2f}: squared error 1626 rewards the average. At β = 3 the code carries nothing and every sample 1627 is the same grey average stroke. 1628 """ 1629 ) 1630 takeaway("β trades rebuild sharpness for a smooth, sampleable code space; latent diffusion keeps it tiny and lets diffusion generate.")