U-Net
U-Net shrinks an image to understand it and grows it back to draw the mask, and the sideways copies are what save the fine detail.
- 14 min read
- 3 reading levels
- Updated
On this page 8
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
U-Net shrinks a picture to understand it, then grows it back to draw the outline. On the way up it keeps peeking at the original.
The analogy
Put a page of small print on a photocopier and shrink it to a quarter size. Now photocopy that copy back up to full size.
The result is a smudge. You can still see there was a paragraph, and roughly where the headings were. You cannot read a word.
That is what happens inside every shrink-and-grow network. Understanding needs shrinking, and detail dies when you shrink.
U-Net's fix is the thing anyone would do: keep the original page beside you while you redraw. Every time you enlarge one step, you glance at the matching original and copy the fine lines back in.
Why it exists
Segmentation asks for a label on every pixel. That is two demands at once, and they fight each other.
To know that a shape is a tumour and not a shadow, you need context. Context means looking at a wide area, and a network gets a wide view by shrinking the picture.
To draw the tumour's edge in the right place, you need precision. Precision means fine detail, which shrinking destroys.
Earlier networks picked one side and lost the other. The outputs looked like blurred blobs that were roughly right and never usable for measurement.
U-Net refuses to choose. It shrinks for understanding, grows back for precision, and carries the detail across sideways.
How it works
original picture
|
[shrink] ─────────── copy ──────────┐
| |
[shrink] ────── copy ──────┐ |
| | |
[shrink] ── copy ──┐ | |
| | | |
smallest picture | | |
understands a lot, | | |
sees no detail | | |
| v v v
[grow back] ──> [grow back] ──> [grow back] ──> the finished maskThe shrinking half is the encoder, meaning the part that builds understanding. The growing half is the decoder, the part that draws the answer.
The sideways arrows are skip connections: copies of the earlier, sharper maps handed straight across. Drawn on paper, the whole thing looks like the letter U, and the name stuck.
Where you have already seen this
- Medical scan tools that outline an organ or a tumour for a radiologist.
- Background removal in video calls, before better methods arrived.
- Satellite maps that shade fields, roads and rooftops.
- Photo apps that select the sky so you can swap it.
U-Net was born in a hospital setting, and the design still shows it. It was built to learn from a handful of labelled images, because nobody had thousands of hand-outlined scans.
What is honestly hard here
The parts of U-Net are ordinary. What is hard is believing how much the sideways copies matter. On the diagram they look like a detail.
They are not a detail. Remove them and the same network learns much more slowly. Its edges never come out right.
Remember this
- Shrinking buys understanding; growing back buys precision.
- Skip connections carry the sharp detail across, so you get both.
- U-Net was designed for small datasets, which is why it still wins on small datasets.
What to learn next
- DeepLab and atrous convolution — the other answer to losing resolution.
- Dice and other segmentation losses — what to optimise once the architecture is in place.
- Autoencoders — the shrink-and-grow idea without the skips.
Developer — Code and libraries.
Setup
pip install torchCPU is fine. The script below trains two networks and takes roughly three minutes on a laptop.
The whole architecture, and the experiment that justifies it
The second network is identical apart from one line: the skip connections are switched off. Everything else, including the seed and the data, is held fixed.
import torch, torch.nn as nn, torch.nn.functional as F
def block(cin, cout):
return nn.Sequential(nn.Conv2d(cin, cout, 3, padding=1), nn.ReLU(),
nn.Conv2d(cout, cout, 3, padding=1), nn.ReLU())
class UNet(nn.Module):
"""A three-level U-Net. use_skips=False keeps every layer but cuts the sideways arrows."""
def __init__(self, use_skips=True, c=8):
super().__init__()
self.use_skips = use_skips
self.enc1, self.enc2, self.enc3 = block(1, c), block(c, 2*c), block(2*c, 4*c)
self.bottleneck = block(4*c, 8*c)
self.up3 = nn.ConvTranspose2d(8*c, 4*c, 2, stride=2)
self.up2 = nn.ConvTranspose2d(4*c, 2*c, 2, stride=2)
self.up1 = nn.ConvTranspose2d(2*c, c, 2, stride=2)
m = 2 if use_skips else 1 # concatenating a skip doubles the input channels
self.dec3, self.dec2, self.dec1 = block(m*4*c, 4*c), block(m*2*c, 2*c), block(m*c, c)
self.head = nn.Conv2d(c, 1, 1) # one logit per pixel
def join(self, x, skip):
return torch.cat([x, skip], dim=1) if self.use_skips else x
def forward(self, x, trace=False):
s1 = self.enc1(x) # full resolution
s2 = self.enc2(F.max_pool2d(s1, 2))
s3 = self.enc3(F.max_pool2d(s2, 2))
b = self.bottleneck(F.max_pool2d(s3, 2)) # everything squeezes through here
d3 = self.dec3(self.join(self.up3(b), s3))
d2 = self.dec2(self.join(self.up2(d3), s2))
d1 = self.dec1(self.join(self.up1(d2), s1))
if trace:
for name, t in [("input", x), ("enc1", s1), ("enc2", s2), ("enc3", s3),
("bottleneck", b), ("dec3", d3), ("dec2", d2), ("dec1", d1)]:
print(f" {name:11s} {tuple(t.shape)}")
return self.head(d1)
# synthetic data: irregular bright regions on a noisy, unevenly lit background
_k = torch.exp(-torch.arange(-4, 5).float()**2 / (2 * 2.5**2))
_k = (_k / _k.sum()).view(1, 1, 1, 9)
def batch(n, gen):
z = torch.randn(n, 1, 32, 32, generator=gen)
z = F.conv2d(F.pad(z, (4, 4, 0, 0), mode="reflect"), _k) # blur across
z = F.conv2d(F.pad(z, (0, 0, 4, 4), mode="reflect"), _k.transpose(2, 3)) # blur down
y = (z > 0).float() # the true mask
img = (0.55 * y + 0.35 * torch.rand(n, 1, 32, 32, generator=gen)
+ torch.linspace(0, 0.25, 32).view(1, 1, 32, 1)).clamp(0, 1)
return img, y
def iou(logits, y):
p = (logits > 0).float()
return ((p * y).sum((1, 2, 3)) / ((p + y) > 0).float().sum((1, 2, 3)).clamp(min=1)).mean().item()
print("shape trace for one 32x32 image:")
UNet().forward(torch.zeros(1, 1, 32, 32), trace=True)
print(" dec3 input = 32 channels from below + 32 channels sideways = 64\n")
val_gen = torch.Generator().manual_seed(123)
xv, yv = batch(64, val_gen)
preds = {}
for use_skips in (True, False):
torch.manual_seed(0)
net = UNet(use_skips)
opt = torch.optim.Adam(net.parameters(), lr=3e-4)
gen = torch.Generator().manual_seed(1)
for step in range(1, 1601):
x, y = batch(8, gen)
loss = F.binary_cross_entropy_with_logits(net(x), y)
opt.zero_grad(); loss.backward(); opt.step()
if step in (400, 1600):
with torch.no_grad():
v = net(xv)
tag = "with skips" if use_skips else "no skips "
print(f"{tag} step {step:>4} params {sum(p.numel() for p in net.parameters()):>7,}"
f" val IoU {iou(v, yv):.3f}")
preds[use_skips] = (v[:1, 0] > 0).float()
print("\nrows 8-17 of one validation image (# = foreground)")
print(f"{'ground truth':<34}{'with skips':<34}{'no skips'}")
for r in range(8, 18):
row = lambda t: "".join("#" if v else "." for v in t[r])
print(f"{row(yv[0, 0]):<34}{row(preds[True][0]):<34}{row(preds[False][0])}")shape trace for one 32x32 image: input (1, 1, 32, 32) enc1 (1, 8, 32, 32) enc2 (1, 16, 16, 16) enc3 (1, 32, 8, 8) bottleneck (1, 64, 4, 4) dec3 (1, 32, 8, 8) dec2 (1, 16, 16, 16) dec1 (1, 8, 32, 32) dec3 input = 32 channels from below + 32 channels sideways = 64 with skips step 400 params 120,681 val IoU 0.931 with skips step 1600 params 120,681 val IoU 0.970 no skips step 400 params 108,585 val IoU 0.487 no skips step 1600 params 108,585 val IoU 0.740 rows 8-17 of one validation image (# = foreground) ground truth with skips no skips ...............########......... ...............########......... ................########........ ...............########......... ...............########......... ................########........ ...............########......... ...............########......... ................########........ .............###########........ .............###########........ ...............#########........ ........#...#############....... ............#############....... ##...........###########........ #########################....... #########################....... #####......#############........ #########################....... #########################....... ########.###############........ #########..##############....... #########################....... ########################........ ######.....##############....... ######.....##############....... ######.#################........ #####......#############........ #####......##############....... #####...###############.........
That run is seeded, on PyTorch 2.5.1 running on CPU. Your numbers should land within a few thousandths; exact floating-point results vary across versions and platforms.
Reading the output
Follow the shape trace first. Height and width fall 32 → 16 → 8 → 4 while channels rise 1 → 8 → 16 → 32 → 64. That is the trade every encoder makes: give up where to learn what. By the bottleneck, one value covers an eight-by-eight patch of the original.
The decoder reverses the sizes but cannot reverse the loss. up3 turns a 4x4 map into an 8x8 map by learning an upsampling. It is inventing the four new values. Without help, it is inventing them from a blurred summary.
Four hundred steps in, the gap is enormous: 0.931 against 0.487. The no-skip network at that point is predicting almost everything as foreground, which is why its score sits near the foreground fraction of the data.
By 1600 steps the no-skip network has caught up part of the way, to 0.740. This is the honest version of the story. Skips do not make the task impossible without them. They make it much faster to learn and much more accurate at the edges.
The printed rows show where the remaining error lives. The no-skip prediction is shifted right by about one pixel and has ragged, wrong boundaries. Its blobs are in the right places. Only the outlines are wrong — which is the entire product in segmentation.
The design decisions worth knowing
Concatenate, do not add. torch.cat stacks the skip beside the upsampled features and lets the next convolution decide how to weigh them. Addition, as in ResNet, forces equal weighting and matching channel counts.
padding=1 with a 3x3 kernel keeps sizes stable, so a skip lines up with its partner without cropping. The 2015 paper used unpadded convolutions, so its 572x572 input produced a 388x388 output, and skips had to be centre-cropped before joining. Almost every modern implementation pads instead.
ConvTranspose2d(c, c, 2, stride=2) doubles the size with learned weights. The common alternative is F.interpolate(x, scale_factor=2) followed by a 3x3 convolution, which avoids the checkerboard artefacts that transposed convolution can produce.
Depth is set by object size, not by fashion. Each pooling step doubles how much of the image one unit sees. If your objects span 200 pixels, you need enough levels to see 200 pixels. Four levels is the classic choice for 512x512 inputs.
Common mistakes
Input size not divisible by the pooling factor. Three pooling steps need sizes divisible by 8. Feed a 30x30 image and the skip is 8x8 while the upsampled map is 6x6, and torch.cat raises a size mismatch. Pad the input up to a multiple, then crop the output back.
Sigmoid applied twice. Use BCEWithLogitsLoss on raw outputs, or BCELoss after a sigmoid — never a sigmoid followed by BCEWithLogitsLoss. The second combination trains slowly and never says why.
Batch norm with batch size 1. High-resolution training forces tiny batches, and BatchNorm2d becomes unstable there. Use GroupNorm or InstanceNorm2d instead.
Evaluating with pixel accuracy. With three per cent foreground, predicting all background scores 97 per cent. Use IoU or Dice, as covered in Dice and other segmentation losses.
Forgetting model.eval() and torch.no_grad() at validation time. Dropout stays on and memory balloons, so your reported score is both wrong and expensive.
Try it yourself
Set c=4 to halve the width, and re-run. The no-skip network degrades sharply while the skip version barely moves. Fine detail is arriving sideways, so the bottleneck does not need to carry it.
What to learn next
- DeepLab and atrous convolution — the other answer to losing resolution.
- Dice and other segmentation losses — what to optimise once the architecture is in place.
- Autoencoders — the shrink-and-grow idea without the skips.
Researcher — Mathematics and papers.
The architecture as published
Ronneberger, Fischer and Brox (2015), U-Net: Convolutional Networks for Biomedical Image Segmentation, MICCAI. Contracting path of repeated $3\times3$ unpadded convolutions with ReLU and $2\times2$ max pooling at stride 2, doubling channels at each downsampling. Expansive path of $2\times2$ up-convolutions halving channels, concatenation with the cropped corresponding feature map, then two $3\times3$ convolutions.
Cropping was necessary because unpadded convolutions shrink each map by 2 pixels per convolution. The published configuration maps a $572\times572$ input to a $388\times388$ output.
Two ideas from that paper are underused today:
Overlap-tile inference. To segment an image larger than memory, predict tiles whose input context extends beyond the output tile, and mirror-pad at the image border. This yields seamless tiling with no border artefacts, and it remains the correct way to run dense prediction on gigapixel inputs.
The weighted loss for touching objects. With per-pixel weights
$$ w(\mathbf{x}) = w_c(\mathbf{x}) + w_0 \cdot \exp!\left(-\frac{(d_1(\mathbf{x}) + d_2(\mathbf{x}))^2}{2\sigma^2}\right) $$
where $w_c$ balances class frequencies, $d_1$ and $d_2$ are distances to the nearest and second-nearest cell borders, $w_0 \approx 10$ and $\sigma \approx 5$ pixels. This places enormous weight on the thin gaps between adjacent cells, forcing the network to keep them open. It is a semantic model producing instance-separable output, and it still outperforms naive approaches on dense touching objects.
Why skips work, beyond "they carry detail"
Three effects are separable.
Information recovery. The bottleneck is an information bottleneck in the literal sense. With output stride $s$, a $H \times W$ input arrives as $\frac{H}{s} \times \frac{W}{s}$, and localisation finer than $s$ pixels must be reconstructed rather than remembered. Skips restore it exactly rather than approximately.
Optimisation. A skip is a short gradient path from the loss to early layers, the same mechanism that makes residual networks trainable (He et al., 2016). The experiment above shows this directly. The gap at 400 steps is far larger than the gap at 1600 steps. A substantial part of the benefit is convergence speed, not final capacity.
Effective ensembling. Following Veit et al. (2016) on residual networks, a U-Net with $L$ skips behaves like a collection of paths of differing depth. Deep paths supply semantics, shallow paths supply localisation, and the decoder convolutions learn the mixture.
Documented variants
| Variant | Change | Why it exists |
|---|---|---|
| 3D U-Net (Çiçek et al., 2016) | volumetric convolutions | CT and MRI are volumes, not slices |
| V-Net (Milletari et al., 2016) | residual blocks, Dice loss | trained directly against the evaluation metric |
| Attention U-Net (Oktay et al., 2018) | gates on the skips | suppresses irrelevant regions carried across |
| U-Net++ (Zhou et al., 2018) | nested dense skips | reduces the semantic gap between encoder and decoder features |
| TransUNet, Swin-UNet (2021) | transformer encoder | global context in the contracting path |
| nnU-Net (Isensee et al., 2021) | no architecture change | automatic configuration of preprocessing, patch size, spacing and augmentation |
nnU-Net is the important entry, and the most-ignored result in the field. It won or matched the state of the art across 23 public biomedical challenges using a plain U-Net. The gain came from systematising the choices around the architecture. Published in Nature Methods, 2021.
The honest reading is uncomfortable for architecture papers: for medical segmentation, preprocessing, patch sampling, augmentation and ensembling explain more variance than the architecture does. Treat any new-architecture claim that has not been compared against a tuned nnU-Net as unmeasured.
Cost
For a symmetric U-Net with base width $c$ and $L$ levels, parameters are dominated by the deepest blocks:
$$ P \approx \sum_{l=0}^{L} \alpha \cdot (2^l c)^2 \cdot k^2 $$
so parameters scale as $4^L$ while activation memory scales the other way, as $H W c \sum_l 4^{-l} \cdot 2^l$. Activations dominate memory at high resolution, which is why gradient checkpointing and mixed precision matter more here than in classification.
Papers
- Ronneberger, Fischer, Brox, U-Net, MICCAI 2015 — arxiv.org/abs/1505.04597
- Çiçek et al., 3D U-Net, MICCAI 2016 — arxiv.org/abs/1606.06650
- Milletari, Navab, Ahmadi, V-Net, 3DV 2016 — arxiv.org/abs/1606.04797
- Oktay et al., Attention U-Net, 2018 — arxiv.org/abs/1804.03999
- Isensee et al., nnU-Net: a self-configuring method for deep learning-based biomedical image segmentation, Nature Methods 2021 — arxiv.org/abs/1809.10486
What to learn next
- DeepLab and atrous convolution — the other answer to losing resolution.
- Dice and other segmentation losses — what to optimise once the architecture is in place.
- Autoencoders — the shrink-and-grow idea without the skips.