Latent Diffusion

Latent Diffusion, U-Nets and Diffusion Transformers

Denoising a 512x512 RGB image means predicting 786,432 numbers per step. Latent diffusion (Rombach et al., the basis of Stable Diffusion 18,251 , 2022) runs the loop in a compressed space instead: a variational autoencoder (VAE) shrinks the image 8 times each way into a 64x64x4 latent, 48 times smaller, and decodes pixels once at the end.

Latent diffusion: the denoiser loops in a small latent space, and the VAE decodes once
Latent diffusion: the denoiser loops in a small latent space, and the VAE decodes once

Stable Diffusion 1.5 and SDXL 1,113 denoise with a U-Net, convolution blocks that shrink and re-expand the latent with skip connections. Peebles and Xie's Diffusion Transformer (DiT, 2022) uses a plain transformer over latent patches, which improves as it grows. Stable Diffusion 3 and FLUX.1 57,408 (12 billion parameters) are transformers trained with rectified flow, a nearly straight path from noise to image that needs fewer steps (FLUX.1 [schnell] uses 1-4). The VAE also explains blurred fine text: each latent position stands for 8x8 pixels.

Latent diffusion's compression: a 512×512 image, its 64×64×4 latent (48× smaller) and the decoded resultHTMLLive
<!doctype html>
<style>
  body { margin: 0; padding: 8px; background: #fafaf7; font: 12px system-ui, sans-serif; color: #263238; }
  canvas { display: block; max-width: 100%; }
</style>
<canvas id="c" width="600" height="330"></canvas>
<p>Each latent position stands for an 8×8 pixel block: fine lettering on a cover does not survive the trip.</p>
<script>
  const c = document.getElementById('c'), ctx = c.getContext('2d');
  // A 128-pixel stand-in for a 512×512 image (each drawn pixel represents 4×4 real pixels)
  const N = 128, img = document.createElement('canvas');
  img.width = img.height = N;
  const g = img.getContext('2d');
  const sky = g.createLinearGradient(0, 0, 0, 80);
  sky.addColorStop(0, '#1f5f8b');  sky.addColorStop(1, '#e09a10');
  g.fillStyle = sky;  g.fillRect(0, 0, N, 80);
  g.fillStyle = '#16435f';  g.fillRect(0, 80, N, 48);
  g.fillStyle = '#fff3e0';  g.fillRect(88, 34, 10, 46);                     // lighthouse
  g.fillStyle = '#b5452f';  g.fillRect(88, 46, 10, 6);  g.fillRect(88, 62, 10, 6);
  g.fillStyle = '#ffe9a8';  g.beginPath();  g.arc(40, 70, 12, 0, 7);  g.fill();   // sun
  g.fillStyle = '#fff';  g.font = 'bold 9px Georgia';  g.fillText('The Quiet Harbor', 14, 118);  // fine text

  // "Encode": average each 8×8 pixel block (2×2 here) into 4 channels
  const L = N / 2, px = g.getImageData(0, 0, N, N).data, latent = [];
  for (let y = 0; y < L; y++) for (let x = 0; x < L; x++) {
    const i = ((y * 2) * N + x * 2) * 4, avg = k => (px[i + k] + px[i + 4 + k] + px[i + N * 4 + k] + px[i + N * 4 + 4 + k]) / 4;
    const [r, gg, b] = [avg(0), avg(1), avg(2)];
    latent.push([r, gg, b, (r + gg + b) / 3]);          // a toy 4th channel: brightness
  }
  // The latent shown as four channel maps, and the "decoded" image scaled back up (blurry)
  const lat = document.createElement('canvas');
  lat.width = lat.height = L;
  const lg = lat.getContext('2d'), ld = lg.createImageData(L, L);
  latent.forEach(([r, gg, b], k) => ld.data.set([r, gg, b, 255], k * 4));
  lg.putImageData(ld, 0, 0);

  ctx.font = '12px system-ui';  ctx.fillStyle = '#263238';
  ctx.imageSmoothingEnabled = false;
  ctx.drawImage(img, 10, 20, 160, 160);
  ctx.fillText('image 512×512×3', 10, 196);  ctx.fillText('786,432 numbers', 10, 212);
  ['R', 'G', 'B', 'luma'].forEach((name, k) => {
    const ch = lg.createImageData(L, L);
    latent.forEach((v, i) => ch.data.set([v[k], v[k], v[k], 255], i * 4));
    const tmp = document.createElement('canvas');  tmp.width = tmp.height = L;
    tmp.getContext('2d').putImageData(ch, 0, 0);
    ctx.drawImage(tmp, 215 + (k % 2) * 62, 40 + Math.floor(k / 2) * 62, 56, 56);
  });
  ctx.fillText('latent 64×64×4', 215, 196);  ctx.fillText('16,384 numbers', 215, 212);
  ctx.imageSmoothingEnabled = true;                     // the decoder upsamples, softly
  ctx.drawImage(lat, 380, 20, 160, 160);
  ctx.fillText('decoded once at the end', 380, 196);  ctx.fillText('(text blurs: 8×8 px per position)', 380, 212);

  // Arrows and labels for the pipeline
  ctx.strokeStyle = '#8d6e63';  ctx.fillStyle = '#8d6e63';
  for (const [x1, x2, label] of [[176, 208, 'VAE encode'], [345, 374, 'VAE decode']]) {
    ctx.beginPath();  ctx.moveTo(x1, 100);  ctx.lineTo(x2, 100);  ctx.stroke();
    ctx.beginPath();  ctx.moveTo(x2, 100);  ctx.lineTo(x2 - 7, 95);  ctx.lineTo(x2 - 7, 105);  ctx.fill();
    ctx.fillText(label, x1 - 6, 90);
  }
  // Size bars: pixels vs latent
  const bars = [['pixel space', 786432, '#b5452f'], ['latent space', 16384, '#3f7d3a']];
  bars.forEach(([name, n, color], k) => {
    const w = 440 * n / 786432, big = w > 300;
    ctx.fillStyle = color;  ctx.fillRect(130, 250 + k * 34, w, 22);
    ctx.fillStyle = '#263238';  ctx.fillText(name, 10, 265 + k * 34);
    ctx.fillStyle = big ? '#fff' : '#263238';
    ctx.fillText(`${n.toLocaleString()} values per denoising step`, big ? 136 : 136 + w, 265 + k * 34);
  });
  ctx.fillText('The U-Net (SD 1.5, SDXL) or diffusion transformer (SD 3, FLUX.1) loops in the small space.', 10, 320);
</script>