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.

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.

Eight strokes, their autoencoder rebuilds from 2 numbers (error 0.07) and their PCA rebuilds from 2 numbers (error 4.42): the autoencoder's are near perfect, PCA's are grey smudges

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.

The plain autoencoder's 200 codes form separate strands, one colour per stroke direction, spread from about -24 to 14, far from the small circle where a standard normal draw usually lands

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.

Left: the autoencoder's codes and 300 random codes from their bounding box, 44 percent marked as junk; right: the worst decoded ones are black blobs and the best are clean strokes

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

Left: the KL penalty is a parabola in the centre mu, with mu = 0.5 costing 0.125; right: in the spread sigma it is zero at sigma = 1, rises steeply towards a pin, with sigma = 0.5 costing 0.318

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.

The VAE's 200 codes packed within about two units of the origin, one strand per stroke direction, each with a small shaded fuzzy region

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.

Left: 24 codes from N(0, I) decoded by the plain autoencoder, many black blobs; right: the same codes decoded by the VAE, clean strokes and a few soft blends

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.

An 11 by 11 grid of codes from -2.2 to 2.2, each decoded by the VAE: neighbouring tiles are similar strokes, and regions of the map hold each direction

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]

Two walks between the same two vertical strokes: the plain autoencoder's passes through a smeared cross 3.88 from any real stroke; the VAE's passes through diagonal strokes, never more than 0.70 from a real one

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.

Left: as beta grows from 0.01 to 3, rebuild error rises, KL falls to zero, and sample distance is lowest near beta 0.3; right: 8 samples per beta, sharp but broken at 0.01, clean at 0.3, softer at 1 and uniform grey smudges at 3

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

on GitHub
   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![Eight strokes, their autoencoder rebuilds from 2 numbers (error 0.07) and their PCA rebuilds from 2 numbers (error 4.42): the autoencoder's are near perfect, PCA's are grey smudges](figures/primer.ml.generative.autoencoders.reconstructions.svg)
 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![The plain autoencoder's 200 codes form separate strands, one colour per stroke direction, spread from about -24 to 14, far from the small circle where a standard normal draw usually lands](figures/primer.ml.generative.autoencoders.ae_codes.svg)
 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![Left: the autoencoder's codes and 300 random codes from their bounding box, 44 percent marked as junk; right: the worst decoded ones are black blobs and the best are clean strokes](figures/primer.ml.generative.autoencoders.holes.svg)
 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![Left: the KL penalty is a parabola in the centre mu, with mu = 0.5 costing 0.125; right: in the spread sigma it is zero at sigma = 1, rises steeply towards a pin, with sigma = 0.5 costing 0.318](figures/primer.ml.generative.autoencoders.kl_penalty.svg)
 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![The VAE's 200 codes packed within about two units of the origin, one strand per stroke direction, each with a small shaded fuzzy region](figures/primer.ml.generative.autoencoders.vae_codes.svg)
 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![Left: 24 codes from N(0, I) decoded by the plain autoencoder, many black blobs; right: the same codes decoded by the VAE, clean strokes and a few soft blends](figures/primer.ml.generative.autoencoders.samples.svg)
 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![An 11 by 11 grid of codes from -2.2 to 2.2, each decoded by the VAE: neighbouring tiles are similar strokes, and regions of the map hold each direction](figures/primer.ml.generative.autoencoders.latent_grid.svg)
 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![Two walks between the same two vertical strokes: the plain autoencoder's passes through a smeared cross 3.88 from any real stroke; the VAE's passes through diagonal strokes, never more than 0.70 from a real one](figures/primer.ml.generative.autoencoders.interpolation.svg)
 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![Left: as beta grows from 0.01 to 3, rebuild error rises, KL falls to zero, and sample distance is lowest near beta 0.3; right: 8 samples per beta, sharp but broken at 0.01, clean at 0.3, softer at 1 and uniform grey smudges at 3](figures/primer.ml.generative.autoencoders.beta_tradeoff.svg)
 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()
Level 3: the code, function by function.
SIDE = 8
STROKE_KINDS = ('horizontal', 'vertical', 'diagonal down', 'diagonal up')
STROKE_WIDTH = 0.6
JUNK_DISTANCE = 1.0
def make_strokes( n: int = 200, seed: int = 0) -> tuple[numpy.ndarray, numpy.ndarray, numpy.ndarray]: on GitHub
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.

def worked_example_line() -> dict: on GitHub
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.

def pca_reconstruction(X: numpy.ndarray, k: int) -> numpy.ndarray: on GitHub
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.

def train_linear_autoencoder( X: numpy.ndarray, n_code: int = 1, steps: int = 2000, lr: float = 0.02, seed: int = 0): on GitHub
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.

def sigmoid(z: numpy.ndarray) -> numpy.ndarray: on GitHub
1101def sigmoid(z: np.ndarray) -> np.ndarray:
1102    # Squashes any number into (0, 1): the decoder's pixels must be valid brightnesses.
1103    return 1 / (1 + np.exp(-z))
class Autoencoder: on GitHub
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.

Autoencoder(n_in: int = 64, n_hidden: int = 32, n_code: int = 2, seed: int = 0) on GitHub
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        }
params
def encode(self, X: numpy.ndarray) -> numpy.ndarray: on GitHub
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).

def decode(self, Z: numpy.ndarray) -> numpy.ndarray: on GitHub
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.

def reconstruct(self, X: numpy.ndarray) -> numpy.ndarray: on GitHub
1150    def reconstruct(self, X: np.ndarray) -> np.ndarray:
1151        return self.decode(self.encode(X))
def loss_and_gradients( self, X: numpy.ndarray, noise: numpy.ndarray | None = None) -> tuple[dict, dict]: on GitHub
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

Reconstruction loss (squared error summed over pixels, averaged over images) and every gradient.

noise is ignored; it is accepted so train can drive this and VAE the same way.

class VAE(Autoencoder): on GitHub
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.

VAE( n_in: int = 64, n_hidden: int = 32, n_code: int = 2, beta: float = 0.3, seed: int = 0) on GitHub
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)
beta
def encode_distribution(self, X: numpy.ndarray) -> tuple[numpy.ndarray, numpy.ndarray]: on GitHub
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).

def encode(self, X: numpy.ndarray) -> numpy.ndarray: on GitHub
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.

def loss_and_gradients( self, X: numpy.ndarray, noise: numpy.ndarray | None = None) -> tuple[dict, dict]: on GitHub
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

def reparameterize( mu: numpy.ndarray, logvar: numpy.ndarray, eps: numpy.ndarray) -> numpy.ndarray: on GitHub
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.

def kl_to_standard_normal(mu: numpy.ndarray, logvar: numpy.ndarray) -> numpy.ndarray: on GitHub
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).

def train( model: Autoencoder, X: numpy.ndarray, steps: int = 1500, lr: float = 0.01, seed: int = 0) -> list[dict]: on GitHub
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.

@lru_cache(maxsize=None)
def trained_autoencoder(steps: int = 1500) -> Autoencoder: on GitHub
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.

@lru_cache(maxsize=None)
def trained_vae( beta: float = 0.3, steps: int = 1500) -> VAE: on GitHub
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.

def reconstruction_error( model: Autoencoder, X: numpy.ndarray) -> float: on GitHub
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.

def distance_to_data(images: numpy.ndarray, X: numpy.ndarray) -> numpy.ndarray: on GitHub
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.

def generate( model: Autoencoder, n: int, seed: int = 0) -> numpy.ndarray: on GitHub
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.

def codes_in_box(codes: numpy.ndarray, n: int, seed: int = 0) -> numpy.ndarray: on GitHub
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.

def interpolate( model: Autoencoder, x_a: numpy.ndarray, x_b: numpy.ndarray, steps: int = 9) -> numpy.ndarray: on GitHub
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).

def largest_step(frames: numpy.ndarray) -> float: on GitHub
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.

def expected_squared_error(guess: float, outcomes) -> float: on GitHub
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.

def beta_sweep(betas=(0.01, 0.03, 0.1, 0.3, 1.0, 3.0)) -> list[dict]: on GitHub
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.

def figures() -> dict: on GitHub
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.

def demo() -> None: on GitHub
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.")