Image and Video Generation

Model Architectures and Training


You have trained a diffusion model for 40,000 steps. The loss curve is beautiful — it fell from 1.02 to 0.043 and flattened. You sample from it and get grey mush with faint dark smudges. Nothing you would call an image.

Nothing is broken. The loss curve simply is not telling you what you assume it is telling you. Here is the arithmetic that explains it. The training target is a standard normal noise vector ε, so a network that outputs zeros for every input scores a mean squared error of exactly E[ε²] = 1.0. That is your floor for "learned literally nothing". Now suppose the network learns only that latents are roughly unit-variance Gaussian and nothing else about images. Under that assumption the best possible prediction is E[ε | zt] = √(1−āt) zt, and its residual error works out to exactly āt. Averaged over a uniform timestep with the standard linear schedule, that is 0.2755.

So the range from 1.0 down to 0.2755 is available to a model that understands nothing but variance. Everything a real image model knows — faces, perspective, the fact that cats have two ears — lives in the narrow band between 0.2755 and roughly 0.03. A drop from 0.045 to 0.043 can be the difference between mush and photographs. Diffusion loss is almost useless as a quality signal. You have to look at samples.

That is the first thing to internalise about training these models. The rest of this lesson is about the network doing the predicting, and the training machinery that makes it converge.

One U-Net step: noisy latent and t go in, noise comes outInput: noisy latent z_t, timestep t, prompt vectorsDown path: resolution halves, channels growBottleneck: self-attentionplus cross-attention on the textUp path: mirrors down, skipconnections restore detailOutput: predicted noise, same shape as the latent
Skip connections are why a U-Net and not a plain stack: the fine detail must survive the trip through the bottleneck.

Why a U-Net and not a plain stack of convolutions

The denoiser has a strange job description: input a tensor of shape 4×64×64, output a tensor of shape 4×64×64, and the output must be spatially aligned with the input at the level of individual cells. A classifier can throw away spatial resolution; this network cannot.

Try the naive design: stack 3×3 convolutions at full resolution and never downsample. Each 3×3 conv grows the receptive field by 2 cells, so to let one output cell see the whole 64-wide latent you need 32 stacked layers — all running at 64×64×320 channels, which is expensive, and none of which get any global context until layer 32. The model cannot tell that the top-left corner is sky while it is deciding what to do with the bottom-right corner, and you get locally plausible texture with no global coherence.

The U-Net solves it with a downsample-then-upsample path. Downsampling is what buys receptive field cheaply. At the 8×8 bottleneck of a 64×64 latent, each cell already summarises an 8×8 region, so a single 3×3 convolution there touches a 24×24 area of the original latent — over a third of the image, from one layer.

LevelLatent resolutionChannels (SD 1.5)What it handles
Input64×64320fine texture, edges, grain
Down 132×32640object parts — an eye, a wheel
Down 216×161280whole objects and their relations
Bottleneck8×81280global composition and scene layout

Downsampling alone would be a disaster, because the fine detail is genuinely gone by the bottleneck and the upsampling path would have to hallucinate it. Skip connections fix that: the feature map at each downsampling level is concatenated onto the matching upsampling level. So the 64×64 output block receives both the coarse decision from the bottleneck ("this region is a face") and the original high-frequency features ("and here are exactly where the edges were"). Cut the skip connections and your outputs go soft and blurry — a distinctive, instantly recognisable failure.

The bottleneck decides what is in the image; the skip connections decide exactly where its edges land. Remove either and you get a characteristic, diagnosable kind of bad output.

Two things make a diffusion U-Net different from the medical-segmentation U-Net the architecture came from. First, every block is conditioned on a timestep, because the right behaviour at t = 900 (invent global structure) is nothing like the right behaviour at t = 50 (sharpen texture and stop). Second, attention layers are interleaved into the convolutional blocks, giving both global spatial reasoning and a channel for the text prompt.

The residual block, and where the timestep enters

The workhorse is a residual block: group normalisation, SiLU activation, 3×3 convolution, twice over, with the input added back at the end. The addition is what allows very deep stacks to train — the gradient has a path straight through.

The timestep is injected between the two convolutions, as a per-channel additive shift:

Python
import torchimport torch.nn as nnclass ResBlock(nn.Module):    def __init__(self, in_ch, out_ch, time_dim):        super().__init__()        self.norm1 = nn.GroupNorm(32, in_ch)        self.conv1 = nn.Conv2d(in_ch, out_ch, 3, padding=1)        # project the timestep embedding to one bias per output channel        self.time_proj = nn.Linear(time_dim, out_ch)        self.norm2 = nn.GroupNorm(32, out_ch)        self.conv2 = nn.Conv2d(out_ch, out_ch, 3, padding=1)        self.skip = (nn.Conv2d(in_ch, out_ch, 1)                     if in_ch != out_ch else nn.Identity())    def forward(self, x, t_emb):        h = self.conv1(torch.nn.functional.silu(self.norm1(x)))        h = h + self.time_proj(torch.nn.functional.silu(t_emb))[:, :, None, None]        h = self.conv2(torch.nn.functional.silu(self.norm2(h)))        return h + self.skip(x)

That [:, :, None, None] is doing real work: it broadcasts one scalar per channel across every spatial position, so the timestep uniformly modulates what each channel means without imposing any spatial pattern. Some architectures go further and use the timestep to predict per-channel scale as well as shift (adaptive group norm), which is stronger conditioning at slightly more cost.

Timestep embedding: why not just feed in the integer

You could pass t = 500 as a single float. It fails badly, for two reasons.

First, scale. Everything else in the network is normalised to roughly unit variance. A value of 500 dominates every activation it touches and destabilises training immediately. You could divide by 1000, but then t = 500 and t = 501 differ by 0.001 — below the noise floor of the surrounding activations, so the network cannot resolve adjacent timesteps at all.

Second, a single number is a poor key. The network needs to condition sharply and differently on hundreds of distinct timesteps. One scalar gives a linear ramp, and the only way to extract "am I near t = 500?" from a ramp is to build a bump detector out of nonlinearities — wasteful, and it must be learned from scratch.

The standard fix borrows sinusoidal positional encoding:

emb(t)2i=sin⁡ ⁣(t100002i/d),emb(t)2i+1=cos⁡ ⁣(t100002i/d)\text{emb}(t)_{2i} = \sin\!\left(\frac{t}{10000^{2i/d}}\right), \qquad \text{emb}(t)_{2i+1} = \cos\!\left(\frac{t}{10000^{2i/d}}\right)

In plain English: represent the timestep as a bank of sine and cosine waves running at geometrically spaced frequencies — a fingerprint rather than a magnitude. Concretely, at d = 320 (SD 1.5's base width) you get 160 frequency pairs. The fastest has period about 6.3 timesteps; the slowest has period about 59,000, far longer than the whole 1,000-step range. So the fast components distinguish t = 500 from t = 501 crisply, while the slow ones encode "we are roughly 50% of the way through" smoothly. Every component sits in [−1, 1], so the scale problem disappears.

Frequency index iPeriod (timesteps)Value at t = 500Encodes
06.3sin = −0.468exact timestep, fine-grained
4062.8sin = −0.262local neighbourhood of ~60 steps
80628sin = −0.959which third of the schedule
15959,317sin = 0.053monotone "how far along", never wraps

The raw sinusoids then pass through a two-layer MLP (320 → 1280 → 1280) to give the embedding that every residual block consumes.

Attention: global consistency, and the prompt

Convolutions are local. Even with the U-Net's receptive field, nothing explicitly forces the left eye and the right eye to be the same colour — they are 20 latent cells apart and only ever meet at the bottleneck, heavily compressed. Self-attention lets every spatial position query every other directly. It is why modern diffusion models produce symmetric faces and consistent lighting across a scene.

The cost is quadratic in the number of positions, which is precisely why you cannot afford it everywhere:

ResolutionPositionsAttention matrix entriesSD 1.5 uses self-attention?
64×644,09616,777,216yes — by far its most expensive layer
32×321,0241,048,576yes
16×1625665,536yes
8×8644,096yes, in the middle block

That 64×64 row is where SD 1.5's memory goes: one 16.8-million-entry matrix per head, per layer, per image. It is why attention slicing and memory-efficient attention kernels matter so much for SD 1.5. SDXL, which works on a 128×128 latent, drops self-attention at its highest resolution entirely and attends only at 64×64 and 32×32 — skipping the top level is what keeps its largest attention matrix affordable.

Cross-attention reuses the same operation with a twist: the queries still come from the image features, but the keys and values come from the 77×768 text embedding produced by CLIP.

Attention(Q,K,V)=softmax ⁣(QK⊤dk)V,Q=WQhimg, K=WKctxt, V=WVctxt\text{Attention}(Q,K,V) = \mathrm{softmax}\!\left(\frac{QK^{\top}}{\sqrt{d_k}}\right)V, \quad Q = W_Q h_{\text{img}},\ K = W_K c_{\text{txt}},\ V = W_V c_{\text{txt}}

In plain English: each spatial position asks which prompt tokens are relevant to it and pulls in a weighted blend of their meanings. At 32×32 that is a 1,024×77 attention matrix — trivially cheap next to the 1,024×1,024 self-attention beside it. Text conditioning is nearly free; global spatial reasoning is what costs money.

In practice self-attention, cross-attention and a feed-forward layer are packaged into a single transformer block that sits after each residual block at the attention resolutions. A SD 1.5 U-Net has roughly 860 million parameters, and the transformer blocks are the majority of them.

Assembling the pieces

Stripped to its skeleton, the whole network is a loop over levels, a bottleneck, and a mirrored loop back up with concatenated skips:

Python
class MiniUNet(nn.Module):    def __init__(self, ch=(128, 256, 512), t_dim=512, ctx=768):        super().__init__()        self.t_mlp = nn.Sequential(nn.Linear(128, t_dim), nn.SiLU(),                                   nn.Linear(t_dim, t_dim))        self.stem = nn.Conv2d(4, ch[0], 3, padding=1)        self.down, self.up, self.skips = nn.ModuleList(), nn.ModuleList(), []        for i in range(len(ch) - 1):            self.down.append(nn.ModuleList([                ResBlock(ch[i], ch[i + 1], t_dim),                CrossAttnBlock(ch[i + 1], ctx) if i > 0 else nn.Identity(),                nn.Conv2d(ch[i + 1], ch[i + 1], 3, stride=2, padding=1)]))        self.mid = ResBlock(ch[-1], ch[-1], t_dim)        for i in reversed(range(len(ch) - 1)):            # in_ch is doubled: upsampled features + concatenated skip            self.up.append(nn.ModuleList([                nn.ConvTranspose2d(ch[i + 1], ch[i + 1], 4, stride=2, padding=1),                ResBlock(ch[i + 1] * 2, ch[i], t_dim)]))        self.out = nn.Conv2d(ch[0], 4, 3, padding=1)   # predict noise, 4 channels    def forward(self, x, t, ctx):        t_emb = self.t_mlp(sinusoidal(t, 128))        h = self.stem(x); skips = []        for res, attn, downsample in self.down:            h = res(h, t_emb)            h = attn(h, ctx) if not isinstance(attn, nn.Identity) else h            skips.append(h); h = downsample(h)        h = self.mid(h, t_emb)        for upsample, res in self.up:            h = upsample(h)            h = res(torch.cat([h, skips.pop()], dim=1), t_emb)        return self.out(h)

Two details there are the ones people get wrong. The skip is concatenated, not added, which is why the upsampling ResBlock takes ch[i+1] * 2 input channels — get that arithmetic wrong and you get a shape error, which is the good outcome; get it subtly wrong with a projection and you silently lose half the skip information. And the final convolution outputs 4 channels, matching the latent, because the target is the noise added to the latent, not an image.

The training loop, in full

Python
import torch, torch.nn.functional as Foptimiser = torch.optim.AdamW(unet.parameters(), lr=1e-4, weight_decay=0.01)scaler = torch.amp.GradScaler("cuda")accum = 8for step, batch in enumerate(dataloader):    with torch.no_grad():        latents = vae.encode(batch["pixels"]).latent_dist.sample()        latents = latents * 0.18215        text_emb = text_encoder(batch["input_ids"])[0]        # 10% conditioning dropout -- this is what makes CFG possible later        drop = torch.rand(len(latents), device=latents.device) < 0.10        text_emb[drop] = empty_embedding    noise = torch.randn_like(latents)    t = torch.randint(0, 1000, (len(latents),), device=latents.device)    noisy = scheduler.add_noise(latents, noise, t)      # closed-form jump    with torch.autocast("cuda", dtype=torch.float16):        pred = unet(noisy, t, encoder_hidden_states=text_emb).sample        loss = F.mse_loss(pred.float(), noise.float()) / accum    scaler.scale(loss).backward()    if (step + 1) % accum == 0:        scaler.unscale_(optimiser)        torch.nn.utils.clip_grad_norm_(unet.parameters(), 1.0)        scaler.step(optimiser); scaler.update()        optimiser.zero_grad()        ema.update(unet)

Five details in there are load-bearing and each is a common bug when omitted.

MistakeWhat actually happensFix
VAE and text encoder left in training mode / not under no_gradgradients flow into frozen modules; memory blows up, VAE slowly degrades.eval(), requires_grad_(False), wrap in no_grad
Forgetting the 0.18215 scale factorlatents have std ~5.5; the schedule's noise levels are wrong for that scale, training silently underperformsmultiply after encoding
Sampling t from a narrow rangemodel never learns the timesteps you skipped; sampling breaks exactly thererandint(0, 1000), uniform, per sample not per batch
No conditioning dropoutno unconditional branch exists, so classifier-free guidance produces garbage at inferencedrop text on ~10% of samples during training
Loss computed in fp16MSE of small values underflows; loss reads as 0 and gradients vanishcast to .float() before the loss

Loss weighting: not all timesteps deserve equal attention

Uniform timestep sampling means uniform loss weighting, and that is measurably suboptimal. Return to the irreducible-loss numbers from the opening. At t = 50 the irreducible loss is 0.971; at t = 800 it is 0.0015. The gradient magnitudes differ by nearly three orders of magnitude, so low-timestep samples dominate the update while contributing least to perceptual quality.

Min-SNR-γ weighting rebalances this. Each sample's loss is scaled by min(SNRt, γ)/SNRt, with γ typically 5:

tSNRMin-SNR-5 weightEffect
5033.500.149heavily down-weighted — nearly-clean inputs stop dominating
1008.710.574partially down-weighted
2001.931.000unchanged
6000.0271.000unchanged — hard, structure-forming timesteps keep full weight

Reported convergence speedups are around 3×. It costs one extra line of code.

The three techniques that decide whether training works

Exponential moving average

Diffusion training is exceptionally noisy: each step sees one random timestep per sample, so consecutive gradient directions barely correlate. The live weights bounce around a good region rather than sitting in it. EMA keeps a shadow copy:

θEMA←d⋅θEMA+(1−d) θ\theta_{\text{EMA}} \leftarrow d \cdot \theta_{\text{EMA}} + (1 - d)\,\theta

With d = 0.9999, the effective averaging window is 1/(1−d) = 10,000 steps, and the half-life is about 6,931 steps — a weight from 6,931 steps ago still contributes half as much as the current one. You sample from the EMA weights and train the live ones.

The quality difference is not subtle. Sampling from live weights instead of EMA weights typically costs 5 to 15 FID points on standard benchmarks; on a small dataset it is the difference between usable and unusable output. Two traps: EMA weights must be saved in your checkpoint (people lose them constantly), and d = 0.9999 needs at least ~30,000 steps to warm up — on a short 3,000-step run it is still mostly your random initialisation, so use d = 0.999 (half-life 693 steps) instead.

If your samples look worse than your loss curve suggests, check that you are sampling from EMA weights before you change anything else.

Gradient accumulation

Diffusion wants large batches precisely because of that gradient noise. If 8 samples is all that fits in memory, run 8 forward/backward passes and step once: effective batch 64 at the memory cost of 8. The only subtlety is that you must divide the loss by the accumulation count, otherwise your effective learning rate is multiplied by it — a silent 8× LR increase that shows up as divergence around step 500 and is very hard to diagnose.

Mixed precision

fp16 halves memory and roughly doubles throughput on tensor-core hardware, but it has a hard floor: the smallest normal fp16 value is about 6.1×10−56.1 \times 10^{-5}. Diffusion gradients frequently sit near 10−710^{-7}, which rounds to exactly zero — those parameters simply stop learning, and nothing in your logs says so.

GradScaler fixes it by multiplying the loss by a large constant (starting at 65,536) before backward, so gradients land in fp16's representable range, then dividing them back out before the optimiser step. It also detects infinities and halves the scale automatically. If you use bf16 instead, you get fp32's exponent range and need no scaler — at the cost of fewer mantissa bits, which is almost always the better trade on Ampere and later.

Hyperparameters that actually move the needle

KnobTypical valueWhat happens if you get it wrong
Learning rate1e-4 (from scratch), 1e-5 to 5e-6 (fine-tune)Too high: loss spikes to NaN, or output collapses to one image. Too low: 3× the training time for the same result.
Effective batch size256–2048Below ~64 the gradient noise overwhelms the signal and the model plateaus early
Warmup steps500–1000Without it, the first large steps from random init frequently produce NaN in the first 50 steps
Gradient clippingmax norm 1.0Without it, one bad batch can destroy hours of training in a single step
EMA decay0.9999 long runs, 0.999 shortToo high on a short run: you sample from noise. Too low: no smoothing benefit.
Base channel width128 (small) to 320 (SD 1.5)Under ~128 the model cannot represent fine texture regardless of training length

Why 1,000 training timesteps but 25 at inference

The network is trained on every integer timestep from 0 to 999, but generation visits perhaps 25 of them. That is legitimate because the timestep embedding is continuous — sinusoids interpolate smoothly, so the network behaves sensibly at timesteps between the ones it happened to see, and it saw them all anyway. The scheduler picks a subset and integrates across the gaps.

The consequence for training is that the schedule you train with is baked into the weights. A model trained on a linear β schedule and sampled with a scheduler configured for a cosine one will produce mismatched noise levels at every step, and the output is a washed-out or over-contrasted mess. Always construct your inference scheduler with from_config(train_scheduler.config) rather than with default arguments — this single mistake accounts for a large share of "my fine-tune generates garbage" reports.

Measuring whether it worked

Since loss will not tell you, you need metrics that compare distributions of images.

FID (Fréchet Inception Distance) embeds real and generated images with an Inception network, fits a Gaussian to each set of features, and measures the distance between them:

FID=∥μr−μg∥2+Tr ⁣(Σr+Σg−2(ΣrΣg)1/2)\text{FID} = \lVert\mu_r - \mu_g\rVert^2 + \mathrm{Tr}\!\left(\Sigma_r + \Sigma_g - 2(\Sigma_r\Sigma_g)^{1/2}\right)

In plain English: do generated images occupy the same region of feature space, with the same spread, as real ones? Lower is better; under 10 is strong, under 5 is state of the art on standard benchmarks. FID is heavily biased by sample count — a score computed on 1,000 images is not comparable to one computed on 50,000, and quoting FID without the count is meaningless. It is also blind to prompt adherence: a model that ignores every prompt and emits gorgeous unrelated photographs scores brilliantly.

CLIP score covers that gap by taking the cosine similarity between the CLIP embedding of the prompt and of the generated image. Typical values run 0.25–0.35. It is the natural companion to FID, and the two trade off against each other — raising guidance scale improves CLIP score and worsens FID, which is why papers report a curve across guidance values rather than a single point.

MetricMeasuresBlind to
FIDrealism and diversity togetherprompt adherence; sensitive to sample count
CLIP scoreprompt adherenceimage quality; saturates on simple prompts
Inception Scoreconfidence and class varietyneeds ImageNet-like classes; largely superseded
Human preferencewhat you actually care aboutslow, expensive, hard to reproduce

A practical compromise is a fixed evaluation set: 20 prompts, 4 fixed seeds each, 80 images regenerated at every checkpoint and viewed as a grid. It catches mode collapse, drift, and overfitting within seconds, which no scalar metric does reliably.

What this means when you train something

Almost nobody trains a diffusion U-Net from scratch. It costs on the order of 150,000 A100-hours for something SD-1.5-class. What you will actually do is fine-tune, and every mechanism above changes character when you do.

Set the learning rate an order of magnitude lower than the from-scratch value — 1e-5, sometimes 5e-6. The pretrained weights already encode enormous knowledge, and 1e-4 will overwrite it within a few hundred steps. That failure has a signature: samples get sharper and more on-style for a while, then every prompt starts returning near-identical images. That is catastrophic forgetting, and the loss curve looks fine throughout.

Watch the trainable surface. Training the whole U-Net on 200 images will memorise them. Training only the cross-attention projections, or attaching low-rank adapters, restricts capacity in a way that produces a style transfer rather than a copy. Freezing the VAE and text encoder is not an optimisation — it is the thing that stops your model degrading in ways you cannot see until you decode.

And build the sample grid before you build anything else. The single most common way a fine-tune is wasted is training for six hours guided by a number that, as the opening arithmetic showed, moves by 0.002 between mush and photographs.