Course Content
Image and Video Generation
4 sections · 7 lessons
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.
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.
| Level | Latent resolution | Channels (SD 1.5) | What it handles |
|---|---|---|---|
| Input | 64×64 | 320 | fine texture, edges, grain |
| Down 1 | 32×32 | 640 | object parts — an eye, a wheel |
| Down 2 | 16×16 | 1280 | whole objects and their relations |
| Bottleneck | 8×8 | 1280 | global 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:
1import torch2import torch.nn as nn34class ResBlock(nn.Module):5 def __init__(self, in_ch, out_ch, time_dim):6 super().__init__()7 self.norm1 = nn.GroupNorm(32, in_ch)8 self.conv1 = nn.Conv2d(in_ch, out_ch, 3, padding=1)9 # project the timestep embedding to one bias per output channel10 self.time_proj = nn.Linear(time_dim, out_ch)11 self.norm2 = nn.GroupNorm(32, out_ch)12 self.conv2 = nn.Conv2d(out_ch, out_ch, 3, padding=1)13 self.skip = (nn.Conv2d(in_ch, out_ch, 1)14 if in_ch != out_ch else nn.Identity())1516 def forward(self, x, t_emb):17 h = self.conv1(torch.nn.functional.silu(self.norm1(x)))18 h = h + self.time_proj(torch.nn.functional.silu(t_emb))[:, :, None, None]19 h = self.conv2(torch.nn.functional.silu(self.norm2(h)))20 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:
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 i | Period (timesteps) | Value at t = 500 | Encodes |
|---|---|---|---|
| 0 | 6.3 | sin = −0.468 | exact timestep, fine-grained |
| 40 | 62.8 | sin = −0.262 | local neighbourhood of ~60 steps |
| 80 | 628 | sin = −0.959 | which third of the schedule |
| 159 | 59,317 | sin = 0.053 | monotone "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:
| Resolution | Positions | Attention matrix entries | SD 1.5 uses self-attention? |
|---|---|---|---|
| 64×64 | 4,096 | 16,777,216 | yes — by far its most expensive layer |
| 32×32 | 1,024 | 1,048,576 | yes |
| 16×16 | 256 | 65,536 | yes |
| 8×8 | 64 | 4,096 | yes, 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.
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:
1class MiniUNet(nn.Module):2 def __init__(self, ch=(128, 256, 512), t_dim=512, ctx=768):3 super().__init__()4 self.t_mlp = nn.Sequential(nn.Linear(128, t_dim), nn.SiLU(),5 nn.Linear(t_dim, t_dim))6 self.stem = nn.Conv2d(4, ch[0], 3, padding=1)7 self.down, self.up, self.skips = nn.ModuleList(), nn.ModuleList(), []8 for i in range(len(ch) - 1):9 self.down.append(nn.ModuleList([10 ResBlock(ch[i], ch[i + 1], t_dim),11 CrossAttnBlock(ch[i + 1], ctx) if i > 0 else nn.Identity(),12 nn.Conv2d(ch[i + 1], ch[i + 1], 3, stride=2, padding=1)]))13 self.mid = ResBlock(ch[-1], ch[-1], t_dim)14 for i in reversed(range(len(ch) - 1)):15 # in_ch is doubled: upsampled features + concatenated skip16 self.up.append(nn.ModuleList([17 nn.ConvTranspose2d(ch[i + 1], ch[i + 1], 4, stride=2, padding=1),18 ResBlock(ch[i + 1] * 2, ch[i], t_dim)]))19 self.out = nn.Conv2d(ch[0], 4, 3, padding=1) # predict noise, 4 channels2021 def forward(self, x, t, ctx):22 t_emb = self.t_mlp(sinusoidal(t, 128))23 h = self.stem(x); skips = []24 for res, attn, downsample in self.down:25 h = res(h, t_emb)26 h = attn(h, ctx) if not isinstance(attn, nn.Identity) else h27 skips.append(h); h = downsample(h)28 h = self.mid(h, t_emb)29 for upsample, res in self.up:30 h = upsample(h)31 h = res(torch.cat([h, skips.pop()], dim=1), t_emb)32 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
1import torch, torch.nn.functional as F23optimiser = torch.optim.AdamW(unet.parameters(), lr=1e-4, weight_decay=0.01)4scaler = torch.amp.GradScaler("cuda")5accum = 867for step, batch in enumerate(dataloader):8 with torch.no_grad():9 latents = vae.encode(batch["pixels"]).latent_dist.sample()10 latents = latents * 0.1821511 text_emb = text_encoder(batch["input_ids"])[0]12 # 10% conditioning dropout -- this is what makes CFG possible later13 drop = torch.rand(len(latents), device=latents.device) < 0.1014 text_emb[drop] = empty_embedding1516 noise = torch.randn_like(latents)17 t = torch.randint(0, 1000, (len(latents),), device=latents.device)18 noisy = scheduler.add_noise(latents, noise, t) # closed-form jump1920 with torch.autocast("cuda", dtype=torch.float16):21 pred = unet(noisy, t, encoder_hidden_states=text_emb).sample22 loss = F.mse_loss(pred.float(), noise.float()) / accum2324 scaler.scale(loss).backward()2526 if (step + 1) % accum == 0:27 scaler.unscale_(optimiser)28 torch.nn.utils.clip_grad_norm_(unet.parameters(), 1.0)29 scaler.step(optimiser); scaler.update()30 optimiser.zero_grad()31 ema.update(unet)Five details in there are load-bearing and each is a common bug when omitted.
| Mistake | What actually happens | Fix |
|---|---|---|
VAE and text encoder left in training mode / not under no_grad | gradients flow into frozen modules; memory blows up, VAE slowly degrades | .eval(), requires_grad_(False), wrap in no_grad |
| Forgetting the 0.18215 scale factor | latents have std ~5.5; the schedule's noise levels are wrong for that scale, training silently underperforms | multiply after encoding |
| Sampling t from a narrow range | model never learns the timesteps you skipped; sampling breaks exactly there | randint(0, 1000), uniform, per sample not per batch |
| No conditioning dropout | no unconditional branch exists, so classifier-free guidance produces garbage at inference | drop text on ~10% of samples during training |
| Loss computed in fp16 | MSE of small values underflows; loss reads as 0 and gradients vanish | cast 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:
| t | SNR | Min-SNR-5 weight | Effect |
|---|---|---|---|
| 50 | 33.50 | 0.149 | heavily down-weighted — nearly-clean inputs stop dominating |
| 100 | 8.71 | 0.574 | partially down-weighted |
| 200 | 1.93 | 1.000 | unchanged |
| 600 | 0.027 | 1.000 | unchanged — 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:
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−5. Diffusion gradients frequently sit near 10−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
| Knob | Typical value | What happens if you get it wrong |
|---|---|---|
| Learning rate | 1e-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 size | 256–2048 | Below ~64 the gradient noise overwhelms the signal and the model plateaus early |
| Warmup steps | 500–1000 | Without it, the first large steps from random init frequently produce NaN in the first 50 steps |
| Gradient clipping | max norm 1.0 | Without it, one bad batch can destroy hours of training in a single step |
| EMA decay | 0.9999 long runs, 0.999 short | Too high on a short run: you sample from noise. Too low: no smoothing benefit. |
| Base channel width | 128 (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:
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.
| Metric | Measures | Blind to |
|---|---|---|
| FID | realism and diversity together | prompt adherence; sensitive to sample count |
| CLIP score | prompt adherence | image quality; saturates on simple prompts |
| Inception Score | confidence and class variety | needs ImageNet-like classes; largely superseded |
| Human preference | what you actually care about | slow, 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.