Segmentation
Be able to tell semantic from instance segmentation and train a simple U-Net.
Prerequisites
- ECNN architectures: LeNet to ResNetrequired
Intuition
Segmentation classifies every pixel. Three variants that are often confused:
| Variant | Answers | Two cars side by side |
|---|---|---|
| Semantic | which class is every pixel? | one connected «car» region |
| Instance | which object does the pixel belong to? | two separate cars |
| Panoptic | both | two cars plus background classes |
U-Net is the architecture that has dominated since 2015, especially in medical image analysis. Its shape is a U:
downsampling upsampling
64 ─────────── skip connection ───────────→ 64
↓ 128 ─────── skip connection ──────→ 128 ↑
↓ 256 ──── skip connection ─→ 256 ↑
↓ 512 ──────────────────→ 512 ↑
The skip connections are the whole point. The downsampling gives context (what is this?) but loses spatial resolution. The skip connections bring back the fine-grained detail from the corresponding level — without them the masks get blurred and the edges imprecise.
Formal
The loss function decides more than the architecture in segmentation, because the classes are nearly always extremely imbalanced — a tumour can be 0.1 % of the pixels.
| Loss | Formula | When |
|---|---|---|
| Cross-entropy | per pixel | balanced classes |
| Weighted cross-entropy | class weights | moderate imbalance |
| Dice | strong imbalance | |
| Dice + CE | the sum | the standard choice in practice |
| Focal | down-weights the easy pixels | extreme imbalance |
| Tversky | Dice with an adjustable FP/FN weight | when misses cost more than false alarms |
Plain cross-entropy on 0.1 % positive pixels gives a model that predicts «background» everywhere and gets 99.9 % pixel accuracy. The Dice loss punishes that directly, since it measures the overlap with the positive class.
Metrics:
| Metric | Formula | Note |
|---|---|---|
| IoU / Jaccard | stricter | |
| Dice / F1 | milder; Dice ≥ IoU always | |
| mIoU | the mean over the classes | the standard for semantic segmentation |
| Hausdorff distance | the largest edge deviation | when the edge is what matters |
Report per class, not just the mean. An mIoU of 0.78 can hide that the most important class is at 0.31.
Practical details that matter a great deal:
| Detail | Why |
|---|---|
| Augmentation with elastic deformation | U-Net's original paper singled this out as decisive for medical images |
| Patch-based training | whole images rarely fit in memory |
| Overlapping patches at inference | otherwise seams show in the mask |
| Class weights from the actual pixel distribution | not guessed |
SAM (Segment Anything) changed the situation in 2023: a promptable model trained on a billion masks that segments arbitrary objects without fine-tuning. For many tasks the right starting point today is to try SAM first and only train your own if that is not enough.
Code
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, bias=False), nn.BatchNorm2d(cout), nn.ReLU(inplace=True),
nn.Conv2d(cout, cout, 3, padding=1, bias=False), nn.BatchNorm2d(cout), nn.ReLU(inplace=True))
class UNet(nn.Module):
def __init__(self, in_channels=3, n_classes=1, base=64):
super().__init__()
self.down1, self.down2 = block(in_channels, base), block(base, base * 2)
self.down3, self.down4 = block(base * 2, base * 4), block(base * 4, base * 8)
self.bottom = block(base * 8, base * 16)
self.up4 = nn.ConvTranspose2d(base * 16, base * 8, 2, stride=2)
self.up3 = nn.ConvTranspose2d(base * 8, base * 4, 2, stride=2)
self.up2 = nn.ConvTranspose2d(base * 4, base * 2, 2, stride=2)
self.up1 = nn.ConvTranspose2d(base * 2, base, 2, stride=2)
self.dec4, self.dec3 = block(base * 16, base * 8), block(base * 8, base * 4)
self.dec2, self.dec1 = block(base * 4, base * 2), block(base * 2, base)
self.out = nn.Conv2d(base, n_classes, 1)
self.pool = nn.MaxPool2d(2)
def forward(self, x):
h1 = self.down1(x)
h2 = self.down2(self.pool(h1))
h3 = self.down3(self.pool(h2))
h4 = self.down4(self.pool(h3))
b = self.bottom(self.pool(h4))
d = self.dec4(torch.cat([self.up4(b), h4], 1)) # the skip connection
d = self.dec3(torch.cat([self.up3(d), h3], 1))
d = self.dec2(torch.cat([self.up2(d), h2], 1))
d = self.dec1(torch.cat([self.up1(d), h1], 1))
return self.out(d)
def dice_loss(logits, target, eps=1.0):
p = torch.sigmoid(logits)
intersection = (p * target).sum(dim=(2, 3))
return (1 - (2 * intersection + eps) / (p.sum(dim=(2, 3)) + target.sum(dim=(2, 3)) + eps)).mean()
def combined(logits, target, w_dice=0.5):
return w_dice * dice_loss(logits, target) + (1 - w_dice) * F.binary_cross_entropy_with_logits(logits, target)
# Why plain cross-entropy fails under strong imbalance
target = torch.zeros(1, 1, 256, 256); target[0, 0, 120:130, 120:130] = 1.0 # 0.15 % positive
all_background = torch.full((1, 1, 256, 256), -10.0) # always predict 0
print("CE :", round(float(F.binary_cross_entropy_with_logits(all_background, target)), 5))
print("Dice:", round(float(dice_loss(all_background, target)), 5))
# CE : 0.00153 ← nearly zero: the model "succeeds" by never finding anything
# Dice: 0.99999 ← punishes it directly
# Measure per class, not just the mean
def iou_per_class(pred, target, n_classes):
out = {}
for k in range(n_classes):
p, m = (pred == k), (target == k)
union = float((p | m).sum())
out[k] = round(float((p & m).sum()) / union, 4) if union else None
return out
# Overlapping patches at inference — otherwise seams show
def segment_large_image(model, image, patch=512, overlap=64):
H, W = image.shape[-2:]
step = patch - overlap
out = torch.zeros(1, 1, H, W)
weight = torch.zeros(1, 1, H, W)
for y in range(0, H, step):
for x in range(0, W, step):
y2, x2 = min(y + patch, H), min(x + patch, W)
with torch.no_grad():
p = torch.sigmoid(model(image[..., y:y2, x:x2]))
out[..., y:y2, x:x2] += p
weight[..., y:y2, x:x2] += 1
return out / weight.clamp_min(1)
The printout in the middle shows why the choice of loss matters more than the architecture: a model that never finds anything gets a cross-entropy of 0.0015 — practically perfect according to that metric.
Mastery means
- Tells semantic, instance and panoptic segmentation apart
- Explains U-Net's skip connections
- Chooses a loss function under imbalance
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — U-Net: Convolutional Networks for Biomedical Image Segmentation — arXiv (open access; licence per article)
- arXiv — Segment Anything — arXiv (open access; licence per article)
- PyTorch — tutorials (BSD-3) — BSD-3-Clause