👁️ Computer Vision · Lecture 14 of 27

Semantic Segmentation: FCN, U-Net and DeepLab

Segmentation labels every pixel. We cover fully convolutional networks, the encoder–decoder U-Net with skip connections, DeepLab's atrous convolutions, loss functions like Dice, and evaluation with IoU.

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:

  1. Replace fully connected layers with convolutions, so the network accepts any input size and outputs a coarse spatial map of class scores.
  2. Upsample the coarse map back to input resolution with learnable transposed convolutions.
  3. 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.

python
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:
$$ \mathcal{L}_{\text{Dice}} = 1 - \frac{2\sum_i p_i g_i + \epsilon}{\sum_i p_i + \sum_i g_i + \epsilon} $$

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.
python
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_pytorch make this easy.
  • Annotation is expensive; consider weak labels (scribbles, boxes), active learning, and foundation models like SAM (covered later) to accelerate labelling.
JA
Written by

Janin A Apurba

B.Sc. in CSE, AUST · Advanced ICT Officer, CNRS-UNHCR. Teaching AI, ML and Deep Learning to the next generation of engineers and researchers.

Keep learning

Related lectures

👁️ Computer Vision

AI in Medical Imaging: Opportunities, Pitfalls and Validation

Deep learning can detect disease in X-rays, retinal scans and pathology slides. We survey modalities and tasks, discuss data and labelling challenges, shortcut learning, rigorous clinical validation, and deployment responsibilities.

Intermediate⏱ 5 min#158
👁️ Computer Vision

Object Detection III: SSD, RetinaNet and the Focal Loss

One-stage detectors face an extreme imbalance between background and objects. We study SSD's multi-scale default boxes, then derive RetinaNet's focal loss, which let one-stage detectors match two-stage accuracy.

Advanced⏱ 5 min#147
👁️ Computer Vision

Instance Segmentation: Mask R-CNN and Beyond

Instance segmentation separates each individual object with its own mask. We study Mask R-CNN's mask branch and RoIAlign, compare instance, semantic and panoptic segmentation, and survey query-based models like Mask2Former.

Advanced⏱ 4 min#149