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.

On this page 7
  1. The short answer
  2. The analogy
  3. What one wipe involves
  4. Why the shape is called a U
  5. Where you have seen this
  6. Remember this
  7. What to learn next

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 times

The 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 size

Those 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

Developer — Code and libraries.

Setup

bash
pip install torch==2.5.1

This 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

diffusion_unet.py
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))
Output
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:

python
xt = abar[t].sqrt() * x0 + (1 - abar[t]).sqrt() * noise

abar 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

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

TargetPredictionNotes
$\boldsymbol{\epsilon}$The noiseDDPM default; poor at very low noise levels
$\mathbf{x}_0$The clean imageBetter 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

What to learn next