Image Generation and Restoration
The denoising U-Net
The network inside a diffusion model does one job, over and over: look at a noisy picture and estimate the noise that was added to it.
- 14 min read
- 3 reading levels
- Updated
Read these first
On this page 7
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
The U-Net looks at a noisy picture and says which part of it is noise.
The analogy
Think about a dusty old photograph in a family album. You wipe it with a cloth. A little dust comes off, and the faces get slightly clearer.
You wipe again. Clearer again. Nobody wipes it clean in one stroke, and nobody needs to.
That is the whole idea. A diffusion model starts from a picture that is nothing but dust. It wipes, hundreds of times, in tiny amounts. What emerges is a picture that was never there.
What one wipe involves
Each wipe is one pass through a network. The network is handed two things.
The noisy picture. And a number saying how far along the process is: heavily noisy, or almost finished.
The network answers one question. Which part of this picture is dust?
noisy picture + step number
|
v
[ U-Net ]
|
v
"this much of it is dust"
|
v
subtract a little of it -> slightly cleaner picture
|
+------> feed back in, hundreds of timesThe step number matters more than you would guess. Wiping a very dusty photo and wiping an almost-clean one are different jobs. One number tells the network which job it is doing.
Why the shape is called a U
The network squeezes the picture down, then builds it back up.
Going down, it loses fine detail but understands the overall scene. Going up, it rebuilds size. The problem is that the fine detail was thrown away on the way down.
So the U-Net cheats, in a good way. At each level on the way up, it gets a copy. That copy shows what the level looked like on the way down.
full size ----------- copied across -----------> full size
| ^
half size -------- copied across --------> half size |
| ^
quarter size ---> the deepest layer ---> quarter sizeThose copies across are called skip connections. They are not a small optimisation. Remove them and the model produces noise instead of pictures. You will see that measured in the next section.
Where you have seen this
- Every text-to-image tool built before the transformer versions arrived.
- Phone "night mode", which cleans up a dark, grainy photo.
- Medical scan tools, where the U-Net shape was invented in the first place.
Remember this
- The U-Net predicts the noise, not the picture.
- It also receives the step number, because early and late steps need different behaviour.
- Skip connections carry fine detail past the bottleneck, and nothing works without them.
What to learn next
- Latent diffusion — running this same network on a compressed image instead of pixels.
- Diffusion models — the process this network sits inside.
- Image segmentation — the task the U-Net was invented for.
Developer — Code and libraries.
Setup
pip install torch==2.5.1This trains two small U-Nets on 16x16 images. It took about 36 seconds on the CPU used here; your machine will differ.
A complete diffusion U-Net, and proof that the skips matter
import torch, torch.nn as nn, math, time
def shapes(n, size=16):
"""Tiny dataset: a white square or a white disc on black, at a random position."""
g = torch.arange(size).float()
yy, xx = torch.meshgrid(g, g, indexing="ij")
c = torch.randint(4, size - 4, (n, 2, 1, 1)).float()
dx, dy, r = xx - c[:, 0], yy - c[:, 1], 3.0
square = ((dx.abs() < r) & (dy.abs() < r)).float()
disc = ((dx ** 2 + dy ** 2) < r * r).float()
pick = (torch.rand(n, 1, 1) < 0.5).float()
return ((pick * square + (1 - pick) * disc) * 2 - 1).unsqueeze(1) # roughly [-1, 1]
def time_embedding(t, dim=32):
"""Same sinusoidal trick transformers use for word positions, applied to the step number."""
half = dim // 2
freqs = torch.exp(-math.log(10000) * torch.arange(half).float() / half)
a = t.float()[:, None] * freqs[None]
return torch.cat([a.sin(), a.cos()], dim=-1)
class Block(nn.Module):
def __init__(self, cin, cout, tdim=32):
super().__init__()
self.c1 = nn.Conv2d(cin, cout, 3, padding=1)
self.c2 = nn.Conv2d(cout, cout, 3, padding=1)
self.t = nn.Linear(tdim, cout) # the step number, injected as a bias
self.n1, self.n2 = nn.GroupNorm(4, cout), nn.GroupNorm(4, cout)
def forward(self, x, temb):
h = torch.nn.functional.silu(self.n1(self.c1(x)))
h = h + self.t(temb)[:, :, None, None]
return torch.nn.functional.silu(self.n2(self.c2(h)))
class UNet(nn.Module):
def __init__(self, skips=True, c=32):
super().__init__()
self.skips = skips
self.d1, self.d2 = Block(1, c), Block(c, c*2)
self.mid = Block(c*2, c*2)
self.u2 = Block(c*2 + (c*2 if skips else 0), c)
self.u1 = Block(c + (c if skips else 0), c)
self.out = nn.Conv2d(c, 1, 1)
self.pool = nn.AvgPool2d(2)
self.up = nn.Upsample(scale_factor=2, mode="nearest")
def forward(self, x, t):
e = time_embedding(t)
h1 = self.d1(x, e) # 16x16
h2 = self.d2(self.pool(h1), e) # 8x8
m = self.mid(self.pool(h2), e) # 4x4
u = self.up(m) # 8x8
u = self.u2(torch.cat([u, h2], 1) if self.skips else u, e)
u = self.up(u) # 16x16
u = self.u1(torch.cat([u, h1], 1) if self.skips else u, e)
return self.out(u)
T = 200
betas = torch.linspace(1e-4, 0.02, T)
abar = torch.cumprod(1 - betas, 0)
def train(skips, steps=400, bs=64, seed=0):
torch.manual_seed(seed)
net = UNet(skips)
opt = torch.optim.Adam(net.parameters(), lr=2e-3)
last = []
for s in range(steps):
x0 = shapes(bs)
t = torch.randint(0, T, (bs,))
noise = torch.randn_like(x0)
a = abar[t][:, None, None, None]
xt = a.sqrt() * x0 + (1 - a).sqrt() * noise # the forward process, in one jump
loss = ((net(xt, t) - noise) ** 2).mean() # predict the noise, not the image
opt.zero_grad(); loss.backward(); opt.step()
if s >= steps - 50: last.append(loss.item())
return net, sum(last) / len(last)
if __name__ == "__main__":
net = UNet(True)
x = shapes(2)
print("parameters:", f"{sum(p.numel() for p in net.parameters()):,}")
print("input", tuple(x.shape), "-> output", tuple(net(x, torch.tensor([5, 199])).shape))
print("time embedding for steps 0, 5, 199 (first 4 numbers):")
for t in (0, 5, 199):
print(f" t={t:3d} {[round(v,3) for v in time_embedding(torch.tensor([t]))[0,:4].tolist()]}")
t0 = time.time()
print("")
print("skip connections mean noise-prediction MSE over the last 50 steps")
nets = {}
for sk in (True, False):
nets[sk], l = train(sk)
print(f" {str(sk):5s} {l:.4f}")
print(f"training seconds: {time.time()-t0:.0f} (varies by machine)")
@torch.no_grad()
def sample(net):
torch.manual_seed(4)
x = torch.randn(1, 1, 16, 16)
for i in reversed(range(T)):
e = net(x, torch.tensor([i]))
mean = (x - betas[i] / (1 - abar[i]).sqrt() * e) / (1 - betas[i]).sqrt()
x = mean + (betas[i].sqrt() * torch.randn_like(x) if i else 0)
return x[0, 0]
for sk in (True, False):
print("")
print("sampled from the model WITH%s skip connections:" % ("" if sk else "OUT"))
for row in sample(nets[sk]):
print(" " + "".join("#" if v > 0 else "." for v in row))parameters: 221,025 input (2, 1, 16, 16) -> output (2, 1, 16, 16) time embedding for steps 0, 5, 199 (first 4 numbers): t= 0 [0.0, 0.0, 0.0, 0.0] t= 5 [-0.959, 0.324, 1.0, 0.777] t=199 [-0.882, -0.929, 0.097, -0.738] skip connections mean noise-prediction MSE over the last 50 steps True 0.0302 False 0.8481 training seconds: 36 (varies by machine) sampled from the model WITH skip connections: ................ ................ .######......... .######......... .######......... .######......... .######......... ................ ................ ................ ................ ................ ................ ................ ................ ................ sampled from the model WITHOUT skip connections: .#...#.#..#..#.. ..#.#.##.....#.# .##..#....###..# .#..##.##...#.#. ...#..##.###..#. #....##.#...#..# #.#..#..#.#..#.. .#....#.##....#. ..#...#....#..#. #.#.#.#.#.#####. .#.###.......##. .....##.##.#.#.. #.##..##.##.#..# ...####..#....## ###..#.......... .....#.##.#.##.#
Exact loss values depend on your PyTorch build. The 25-fold gap, and the difference between the two sampled pictures, is the reproducible part.
Reading the output
Input shape equals output shape. (2, 1, 16, 16) in, (2, 1, 16, 16) out. The U-Net predicts one noise value per pixel, so it is a dense prediction network, not a classifier.
The time embedding at step 0 is all zeros. That is correct and it is a real gotcha. sin(0) = 0 fills the first half, and cos(0) = 1 fills the second. That is why only the first four numbers are zero. Steps 5 and 199 give plainly different vectors. That is the entire point: one network, many jobs, told apart by this vector.
Skip connections change the loss by a factor of 25. 0.0302 against 0.8481. Noise has variance close to 1, so a loss near 0.85 means the network is barely predicting anything. Removing the skips did not make the model slightly worse. It broke it.
The sampled pictures make it concrete. With skips, a clean square appears in the top-left region. Without skips, the sample is indistinguishable from static. This is why every diffusion U-Net you will ever read has them.
The forward process, done in one line
The loop above never adds noise step by step. It jumps straight to step t:
xt = abar[t].sqrt() * x0 + (1 - abar[t]).sqrt() * noiseabar is the cumulative product of 1 - beta. Because the sum of Gaussians is Gaussian, t steps of small noise collapse into a single closed-form expression. That is why training can sample a random t per example instead of simulating a chain. It is what makes diffusion training affordable at all.
What a real diffusion U-Net adds
The block above is the skeleton. Production models (Stable Diffusion 1.5 and SDXL, for instance) add:
- Residual connections inside each block, so the block learns a correction rather than a whole transform.
- Self-attention at the lower resolutions, usually at 16x16 and 8x8. Attention at full resolution is unaffordable.
- Cross-attention to a text embedding, which is how a prompt gets in. The text is the keys and values, the image features are the queries.
- More channels at lower resolution. 320 to 640 to 1280 in SD 1.5, mirroring the trade every convolutional network makes.
Common mistakes
Predicting the image instead of the noise. Both are valid parameterisations. Noise prediction has roughly constant target scale across timesteps, which conditions the optimisation far better. Ho et al. (2020) found it works substantially better in practice.
Forgetting to feed the timestep. A U-Net without it still trains and the loss still falls. Sampling then produces mush, because the network has learned an average over all noise levels.
Concatenating skips in the wrong order. torch.cat([up, skip], dim=1) must match the channel count declared in the next block. Getting the order wrong is silent when the counts happen to match.
Using batch-norm. Diffusion U-Nets use group normalisation. Batch-norm mixes statistics across a batch containing many different timesteps, which is meaningless here.
Training on data outside [-1, 1]. The noise schedule assumes it. Feed [0, 255] and the signal drowns the noise at every step.
Try it yourself
Set T = 20 instead of 200, retrain, and sample. The steps become large, the model has less to learn, and the shapes come out ragged. Then try c=8 and watch how small a U-Net can still be and work.
What to learn next
- Latent diffusion — running this same network on a compressed image instead of pixels.
- Diffusion models — the process this network sits inside.
- Image segmentation — the task the U-Net was invented for.
Researcher — Mathematics and papers.
What the network is estimating
Ho et al. (2020), Denoising Diffusion Probabilistic Models, define a fixed forward process:
$$ q(\mathbf{x}t \mid \mathbf{x}{t-1}) = \mathcal{N}!\left(\mathbf{x}_t; \sqrt{1 - \beta_t}\,\mathbf{x}_{t-1},\; \beta_t \mathbf{I}\right) $$
$\beta_t$ is the variance added at step $t$. With $\alpha_t = 1 - \beta_t$ and $\bar{\alpha}t = \prod{s=1}^{t} \alpha_s$, the marginal has a closed form:
$$ q(\mathbf{x}_t \mid \mathbf{x}_0) = \mathcal{N}!\left(\mathbf{x}_t; \sqrt{\bar{\alpha}_t}\,\mathbf{x}_0,\; (1 - \bar{\alpha}_t)\mathbf{I}\right) $$
The simplified training objective is a plain regression:
$$ L_{\text{simple}} = \mathbb{E}_{t, \mathbf{x}0, \boldsymbol{\epsilon}} \left[\left| \boldsymbol{\epsilon} - \boldsymbol{\epsilon}\theta!\left(\sqrt{\bar{\alpha}_t}\mathbf{x}_0 + \sqrt{1 - \bar{\alpha}_t}\,\boldsymbol{\epsilon},\; t\right) \right|^2\right] $$
$\boldsymbol{\epsilon} \sim \mathcal{N}(0, \mathbf{I})$ is the sampled noise and $\boldsymbol{\epsilon}_\theta$ the network. This is the variational bound with the per-timestep weights dropped, which the authors found empirically better.
The score connection
Song and Ermon (2019) and Song et al. (2021) show $\boldsymbol{\epsilon}_\theta$ is a rescaled score estimator:
$$ \nabla_{\mathbf{x}_t} \log q(\mathbf{x}t) = -\frac{\boldsymbol{\epsilon}\theta(\mathbf{x}_t, t)}{\sqrt{1 - \bar{\alpha}_t}} $$
So the U-Net estimates the gradient of the log density of the noised data. This unifies DDPM with score matching and with the continuous-time SDE formulation. It is why every sampler in the next lesson can be derived as an ODE or SDE solver.
Parameterisations
| Target | Prediction | Notes |
|---|---|---|
| $\boldsymbol{\epsilon}$ | The noise | DDPM default; poor at very low noise levels |
| $\mathbf{x}_0$ | The clean image | Better at high noise; unstable at low |
| $\mathbf{v}$ | $\sqrt{\bar{\alpha}_t}\boldsymbol{\epsilon} - \sqrt{1-\bar{\alpha}_t}\mathbf{x}_0$ | Salimans and Ho (2022); stable across the whole range |
$\mathbf{v}$-prediction is what distillation and most zero-terminal-SNR schedules use. $\boldsymbol{\epsilon}$-prediction is degenerate at $\bar{\alpha}_t = 0$. With pure noise as input, predicting the input is a perfect solution that carries no information.
Architecture, precisely
The DDPM backbone is the PixelCNN++ / Wide ResNet U-Net with:
- Residual blocks at each resolution, group normalisation, SiLU activations.
- Sinusoidal timestep embedding through a two-layer MLP, added as a per-channel bias inside every residual block.
- Self-attention at 16x16 resolution.
- Downsampling by strided convolution, upsampling by nearest-neighbour plus convolution.
Dhariwal and Nichol (2021), Diffusion Models Beat GANs on Image Synthesis, ablate this thoroughly. Their findings that transferred are these. Attention at multiple resolutions, 32, 16 and 8. More heads with fewer channels each. BigGAN-style residual blocks for resampling. And adaptive group normalisation, where the timestep and class embedding produce the scale and shift of each group-norm layer.
The honest architectural update
The U-Net is no longer the frontier backbone. Peebles and Xie (2023), Scalable Diffusion Models with Transformers (DiT), replace it with a plain transformer over latent patches. They show cleaner scaling with compute.
Stable Diffusion 3 (Esser et al., 2024) introduced MMDiT. It is a multimodal diffusion transformer. Image and text tokens have separate weights, joined in a shared attention operation. FLUX uses a hybrid of single-stream and double-stream MMDiT blocks. SD 3.5 uses MMDiT-X, with QK-normalisation and dual attention in the early layers. MMDiT is now the standard backbone in current open text-to-image models.
The U-Net is still worth understanding for three reasons. Stable Diffusion 1.5 and SDXL remain the most fine-tuned models in the open ecosystem, and both are U-Nets. ControlNet, LoRA adapters and the entire community tooling stack were designed against U-Net block structure. And the U-Net is still the standard backbone for dense restoration tasks. Super-resolution and denoising both have input and output on aligned pixel grids.
Papers
- Ronneberger et al., U-Net: Convolutional Networks for Biomedical Image Segmentation, 2015 — arxiv.org/abs/1505.04597
- Ho et al., Denoising Diffusion Probabilistic Models, 2020 — arxiv.org/abs/2006.11239
- Song et al., Score-Based Generative Modeling through Stochastic Differential Equations, 2021 — arxiv.org/abs/2011.13456
- Nichol and Dhariwal, Improved Denoising Diffusion Probabilistic Models, 2021 — arxiv.org/abs/2102.09672
- Dhariwal and Nichol, Diffusion Models Beat GANs on Image Synthesis, 2021 — arxiv.org/abs/2105.05233
- Salimans and Ho, Progressive Distillation for Fast Sampling of Diffusion Models, 2022 — arxiv.org/abs/2202.00512
- Peebles and Xie, Scalable Diffusion Models with Transformers, 2023 — arxiv.org/abs/2212.09748
- Esser et al., Scaling Rectified Flow Transformers for High-Resolution Image Synthesis, 2024 — arxiv.org/abs/2403.03206
What to learn next
- Latent diffusion — running this same network on a compressed image instead of pixels.
- Diffusion models — the process this network sits inside.
- Image segmentation — the task the U-Net was invented for.