Classification labels an image; detection draws boxes; semantic segmentation assigns a class to every pixel. It outlines tumours in MRI scans, maps flooded areas and informal settlements from satellite images, identifies drivable road for autonomous vehicles and separates crops from weeds. It is the most detailed form of image understanding among the classic tasks.
The task#
Input: an image $H \times W \times 3$. Output: a label map $H \times W$ with one of $K$ classes per pixel (or $K$ probability maps). "Semantic" segmentation does not separate instances: two adjacent people are both simply "person". (Instance segmentation, next lecture, separates them.)
Fully Convolutional Networks (FCN)#
Long, Shelhamer and Darrell (2015) observed that a classification CNN becomes a dense predictor if you:
- Replace fully connected layers with convolutions, so the network accepts any input size and outputs a coarse spatial map of class scores.
- Upsample the coarse map back to input resolution with learnable transposed convolutions.
- Add skip connections from earlier, higher-resolution layers (FCN-16s, FCN-8s) to recover fine detail.
This end-to-end, pixels-to-pixels formulation founded modern segmentation.
U-Net#
Ronneberger, Fischer and Brox (2015) designed U-Net for biomedical images with very few training examples. Its symmetric encoder–decoder has a characteristic U shape:
- Encoder (contracting path): repeated conv blocks and downsampling capture context ("what").
- Decoder (expanding path): upsampling and conv blocks recover resolution ("where").
- Skip connections: at each level, encoder feature maps are concatenated with decoder feature maps, giving the decoder precise spatial detail.
U-Net trained well from only dozens of annotated images with heavy augmentation (elastic deformations) and became the default architecture for medical segmentation — and, later, the backbone of diffusion image generators.
import torch
import torch.nn as nn
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_ch=3, n_classes=2, base=32):
super().__init__()
c = [base, base * 2, base * 4, base * 8]
self.enc = nn.ModuleList([block(in_ch, c[0]), block(c[0], c[1]), block(c[1], c[2])])
self.pool = nn.MaxPool2d(2)
self.mid = block(c[2], c[3])
self.up = nn.ModuleList([nn.ConvTranspose2d(c[3], c[2], 2, 2), nn.ConvTranspose2d(c[2], c[1], 2, 2),
nn.ConvTranspose2d(c[1], c[0], 2, 2)])
self.dec = nn.ModuleList([block(c[3], c[2]), block(c[2], c[1]), block(c[1], c[0])])
self.head = nn.Conv2d(c[0], n_classes, 1)
def forward(self, x):
skips = []
for e in self.enc:
x = e(x); skips.append(x); x = self.pool(x)
x = self.mid(x)
for up, d, s in zip(self.up, self.dec, reversed(skips)):
x = d(torch.cat([up(x), s], dim=1)) # skip connection by concatenation
return self.head(x) # (N, n_classes, H, W) logits
print(UNet()(torch.randn(1, 3, 128, 128)).shape)DeepLab and atrous convolution#
Downsampling loses detail, but context needs large receptive fields. The DeepLab family (Chen et al., 2015–2018) used:
- Atrous (dilated) convolutions to enlarge receptive fields while keeping feature maps at higher resolution;
- Atrous Spatial Pyramid Pooling (ASPP): parallel dilated convolutions at several rates plus global pooling, capturing multi-scale context;
- DeepLabv3+: adds a lightweight decoder for sharper boundaries.
Loss functions#
- Pixel-wise cross-entropy — the default.
- Weighted cross-entropy — up-weight rare classes (small lesions, thin roads).
- Dice loss — directly optimises overlap, robust to foreground–background imbalance:
where $p_i$ are predicted probabilities and $g_i$ ground-truth labels. Combining cross-entropy and Dice is common in medical imaging.
- Focal loss and boundary losses for hard pixels and edges.
def dice_loss(logits, target, eps=1.0):
p = torch.sigmoid(logits).flatten(1); g = target.flatten(1).float()
return 1 - ((2 * (p * g).sum(1) + eps) / (p.sum(1) + g.sum(1) + eps)).mean()Evaluation#
- Pixel accuracy — misleading when background dominates.
- IoU (Jaccard) per class and mean IoU (mIoU) — the standard metric.
- Dice coefficient (equivalent to F1 over pixels) — standard in medicine.
- Boundary metrics (e.g. Hausdorff distance) when contours matter clinically.
Practical tips#
- High-resolution images (satellite, pathology slides) are processed in tiles with overlap; predictions are stitched and blended.
- Use pretrained encoders (e.g. a ResNet or EfficientNet encoder in a U-Net) — libraries like
segmentation_models_pytorchmake this easy. - Annotation is expensive; consider weak labels (scribbles, boxes), active learning, and foundation models like SAM (covered later) to accelerate labelling.