Skip to content
AI-grafen
FAI engineeringComputer vision· about 90 min· fast-moving, sources checked often· verified 2026-09-21· EN

Segmentation

Be able to tell semantic from instance segmentation and train a simple U-Net.

Prerequisites

Intuition

Segmentation classifies every pixel. Three variants that are often confused:

VariantAnswersTwo cars side by side
Semanticwhich class is every pixel?one connected «car» region
Instancewhich object does the pixel belong to?two separate cars
Panopticbothtwo 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.

LossFormulaWhen
Cross-entropyper pixelbalanced classes
Weighted cross-entropyclass weightsmoderate imbalance
Dice1−2∥X∩Y∥∥X∥+∥Y∥1 - \frac{2\|X\cap Y\|}{\|X\|+\|Y\|}strong imbalance
Dice + CEthe sumthe standard choice in practice
Focaldown-weights the easy pixelsextreme imbalance
TverskyDice with an adjustable FP/FN weightwhen 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:

MetricFormulaNote
IoU / Jaccard∥X∩Y∥∥X∪Y∥\frac{\|X\cap Y\|}{\|X\cup Y\|}stricter
Dice / F12∥X∩Y∥∥X∥+∥Y∥\frac{2\|X\cap Y\|}{\|X\|+\|Y\|}milder; Dice ≥ IoU always
mIoUthe mean over the classesthe standard for semantic segmentation
Hausdorff distancethe largest edge deviationwhen 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:

DetailWhy
Augmentation with elastic deformationU-Net's original paper singled this out as decisive for medical images
Patch-based trainingwhole images rarely fit in memory
Overlapping patches at inferenceotherwise seams show in the mask
Class weights from the actual pixel distributionnot 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

All the sources and licences