Lesson 42: Semantic Segmentation

Classification labels a whole image. Detection (Lessons 40-41) labels a handful of boxes. Semantic segmentation goes one step further: label every pixel with a class/category. The original deep-learning method is the fully convolutional network (FCN) (Long, Shelhamer & Darrell, 2015★), which replaces a classifier's fully-connected layers with convolutions, then upsamples the result back to the input resolution. The output is a dense, per-pixel prediction instead of a single label for the whole image. This lesson builds that idea from scratch, checks it against a real pretrained FCN, and then adds one specific enhancement to it — skip connections.

In [1]:
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import matplotlib.pyplot as plt

A 3-class pixel labeling task

Create an image with several small circles and squares scattered on a background. The target is a full per-pixel class map: 0 = background, 1 = circle, 2 = square.

In [2]:
def make_scene(rng, size=32, n_shapes=5):
    scene = np.zeros((size, size), dtype=np.float32)
    mask = np.zeros((size, size), dtype=np.int64)  # 0=background, 1=circle, 2=square
    for _ in range(n_shapes):
        shape_type = rng.choice([1, 2])
        r = rng.integers(2, 4)
        cx, cy = rng.integers(r, size - r), rng.integers(r, size - r)
        yy, xx = np.mgrid[0:size, 0:size]
        if shape_type == 1:
            m = ((xx - cx) ** 2 + (yy - cy) ** 2) <= r ** 2
        else:
            m = (np.abs(xx - cx) <= r) & (np.abs(yy - cy) <= r)
        scene[m] = 1.0
        mask[m] = shape_type
    scene = np.clip(scene + rng.normal(0, 0.05, scene.shape), 0, 1).astype(np.float32)
    return scene, mask

rng = np.random.default_rng(13)
N = 300
scenes, masks = [], []
for _ in range(N):
    s, m = make_scene(rng)
    scenes.append(s); masks.append(m)
scenes = np.array(scenes, dtype=np.float32)
masks = np.array(masks, dtype=np.int64)

split = int(0.85 * N)
Xtr, Mtr = scenes[:split], masks[:split]
Xte, Mte = scenes[split:], masks[split:]

class_fracs = (np.bincount(Mtr.ravel()) / Mtr.size).round(3).tolist()
print('class pixel fractions (train):', dict(zip(['background', 'circle', 'square'], class_fracs)))

fig, axes = plt.subplots(2, 4, figsize=(9, 4.5))
for i in range(4):
    axes[0, i].imshow(Xtr[i], cmap='gray'); axes[0, i].axis('off')
    axes[1, i].imshow(Mtr[i], cmap='viridis', vmin=0, vmax=2); axes[1, i].axis('off')
axes[0, 0].set_title('input', fontsize=9, loc='left')
axes[1, 0].set_title('per-pixel ground truth label', fontsize=9, loc='left')
plt.show()
class pixel fractions (train): {'background': 0.866, 'circle': 0.046, 'square': 0.088}
No description has been provided for this image

A minimal FCN-style network

A segmentation network outputs a class-probability vector for every pixel, so its output is the same spatial size as the input. The encoder is an ordinary CNN, downsampling twice via max pooling (Lesson 34) to build up semantically rich features with a wide receptive field (Lesson 11's pyramid). The decoder then upsamples that coarse bottleneck straight back to full resolution. That's the entire FCN recipe: just downsample then upsample.

In [3]:
class FCNTiny(nn.Module):
    def __init__(self, n_classes=3):
        super().__init__()
        self.enc1 = nn.Sequential(nn.Conv2d(1, 16, 3, padding=1), nn.ReLU())
        self.enc2 = nn.Sequential(nn.Conv2d(16, 32, 3, padding=1), nn.ReLU())
        self.enc3 = nn.Sequential(nn.Conv2d(32, 64, 3, padding=1), nn.ReLU())
        self.pool = nn.MaxPool2d(2)
        self.up = nn.Upsample(scale_factor=2, mode='nearest')
        self.dec2 = nn.Sequential(nn.Conv2d(64, 32, 3, padding=1), nn.ReLU())
        self.dec1 = nn.Sequential(nn.Conv2d(32, 16, 3, padding=1), nn.ReLU())
        self.out = nn.Conv2d(16, n_classes, 1)

    def forward(self, x):
        f1 = self.enc1(x)                # (B,16,H,W)
        f2 = self.enc2(self.pool(f1))    # (B,32,H/2,W/2)
        f3 = self.enc3(self.pool(f2))    # (B,64,H/4,W/4)  -- the bottleneck
        d2 = self.dec2(self.up(f3))      # upsample only -- no fusion with f2
        d1 = self.dec1(self.up(d2))      # upsample only -- no fusion with f1
        return self.out(d1)

def mean_iou(preds, targets, n_classes=3):
    ious = []
    for c in range(n_classes):
        p, t = preds == c, targets == c
        inter, union = (p & t).sum(), (p | t).sum()
        ious.append(inter / union if union > 0 else float('nan'))
    return ious

torch.manual_seed(0)
model_fcn = FCNTiny()
opt = torch.optim.Adam(model_fcn.parameters(), lr=0.01)
Xt = torch.tensor(Xtr).unsqueeze(1); Mt = torch.tensor(Mtr)
for _ in range(200):
    opt.zero_grad()
    loss = F.cross_entropy(model_fcn(Xt), Mt)
    loss.backward()
    opt.step()

with torch.no_grad():
    preds_fcn = model_fcn(torch.tensor(Xte).unsqueeze(1)).argmax(1).numpy()

pixel_acc_fcn = (preds_fcn == Mte).mean()
ious_fcn = mean_iou(preds_fcn, Mte)
print(f'pixel accuracy: {pixel_acc_fcn:.1%}')
for name, iou in zip(['background', 'circle', 'square'], ious_fcn):
    print(f'  {name:>10} IoU: {iou:.3f}')
print(f'mean IoU: {np.mean(ious_fcn):.3f}')
pixel accuracy: 98.1%
  background IoU: 0.984
      circle IoU: 0.736
      square IoU: 0.891
mean IoU: 0.870

In practice: FCN on a real photo

torchvision ships FCN (Long et al., 2015★), pretrained on COCO images (labeled with the 21 Pascal VOC categories — 20 object classes plus background). This is architecturally the same recipe as FCNTiny above: encoder, then upsampling.

In [4]:
import torchvision
import cv2
import matplotlib.patches as mpatches

weights = torchvision.models.segmentation.FCN_ResNet50_Weights.COCO_WITH_VOC_LABELS_V1
fcn = torchvision.models.segmentation.fcn_resnet50(weights=weights)
fcn.eval()  # frozen, no training at all
voc_classes = weights.meta['categories']

photo = cv2.imread('../img/bicycle_train.jpg')
photo_rgb = cv2.cvtColor(photo, cv2.COLOR_BGR2RGB)

preprocess = weights.transforms()
photo_tensor = torch.tensor(photo_rgb / 255.0, dtype=torch.float32).permute(2, 0, 1)
batch = preprocess(photo_tensor).unsqueeze(0)

with torch.no_grad():
    real_out = fcn(batch)['out']
real_pred = real_out.argmax(1)[0].numpy()
# the model's internal resolution differs from the photo's; resize the label map back up to match
real_pred_full = cv2.resize(real_pred.astype(np.uint8), (photo_rgb.shape[1], photo_rgb.shape[0]),
                             interpolation=cv2.INTER_NEAREST)

found_ids = sorted(c for c in np.unique(real_pred_full) if c != 0)
colors = plt.cm.tab10(np.linspace(0, 1, 10))

label_rgb = np.full((*real_pred_full.shape, 3), 255, dtype=np.uint8)  # background = white
for i, cls_id in enumerate(found_ids):
    color_rgb = (np.array(colors[i][:3]) * 255).astype(np.uint8)
    label_rgb[real_pred_full == cls_id] = color_rgb

fig, axes = plt.subplots(1, 2, figsize=(12, 4.5))
axes[0].imshow(photo_rgb); axes[0].set_title('input'); axes[0].axis('off')
axes[1].imshow(label_rgb); axes[1].set_title('FCN per-pixel labels'); axes[1].axis('off')
legend_handles = [mpatches.Patch(color=colors[i], label=voc_classes[c]) for i, c in enumerate(found_ids)]
axes[1].legend(handles=legend_handles, loc='upper right', fontsize=8)
plt.tight_layout()
plt.show()
No description has been provided for this image

Photo by Felix Hanspach on Unsplash

The cyclist is cleanly labeled person, with the bicycle's wheels correctly labeled bicycle just below — both real, correctly-placed detections. The tram is where it gets interesting: the model splits it between train (the upper, windowed section) and bus (the lower body) rather than picking one label for the whole vehicle. That's not really a mistake so much as an honest reflection of VOC's limited vocabulary — a tram is a genuinely ambiguous case between "train" and "bus" for a model that was never given a "tram" class to choose from, and the two guesses roughly track a real visual seam in the vehicle (windows vs. body) rather than being random. A second, smaller person blob on the right edge catches the blurred pedestrian in the background.

Adding skip connections: U-Net

FCNTiny's decoder only sees the pooled, upsampled bottleneck features. After two rounds of pooling, that bottleneck has just a quarter of the input's spatial resolution, and small or closely-packed shapes can blur together or lose their boundaries entirely at that resolution. U-Net (Ronneberger, Fischer & Brox, 2015★) fixes this with skip connections: at each decoder stage, it concatenates the encoder's feature map from the matching resolution, preserving the fine-grained details that would otherwise be lost during pooling. This is similar in spirit to Lesson 37's residual connections, but the purpose and mechanism differ: U-Net concatenates features across the encoder-decoder boundary, rather than adding a residual update within a single network stack.

In [5]:
class UNetTiny(nn.Module):
    def __init__(self, n_classes=3):
        super().__init__()
        self.enc1 = nn.Sequential(nn.Conv2d(1, 16, 3, padding=1), nn.ReLU())
        self.enc2 = nn.Sequential(nn.Conv2d(16, 32, 3, padding=1), nn.ReLU())
        self.enc3 = nn.Sequential(nn.Conv2d(32, 64, 3, padding=1), nn.ReLU())
        self.pool = nn.MaxPool2d(2)
        self.up = nn.Upsample(scale_factor=2, mode='nearest')
        self.dec2 = nn.Sequential(nn.Conv2d(64 + 32, 32, 3, padding=1), nn.ReLU())
        self.dec1 = nn.Sequential(nn.Conv2d(32 + 16, 16, 3, padding=1), nn.ReLU())
        self.out = nn.Conv2d(16, n_classes, 1)

    def forward(self, x):
        f1 = self.enc1(x)                                     # (B,16,H,W)
        f2 = self.enc2(self.pool(f1))                          # (B,32,H/2,W/2)
        f3 = self.enc3(self.pool(f2))                          # (B,64,H/4,W/4)
        d2 = self.dec2(torch.cat([self.up(f3), f2], dim=1))    # skip from f2
        d1 = self.dec1(torch.cat([self.up(d2), f1], dim=1))    # skip from f1
        return self.out(d1)

torch.manual_seed(0)
model_unet = UNetTiny()
opt2 = torch.optim.Adam(model_unet.parameters(), lr=0.01)
for _ in range(200):
    opt2.zero_grad()
    loss = F.cross_entropy(model_unet(Xt), Mt)
    loss.backward()
    opt2.step()

with torch.no_grad():
    preds_unet = model_unet(torch.tensor(Xte).unsqueeze(1)).argmax(1).numpy()

pixel_acc_unet = (preds_unet == Mte).mean()
ious_unet = mean_iou(preds_unet, Mte)

print(f'{"":>18} {"pixel acc":>10} {"mean IoU":>10}')
print(f'{"FCN (no skip)":>18} {pixel_acc_fcn:>10.1%} {np.mean(ious_fcn):>10.3f}')
print(f'{"U-Net (skip)":>18} {pixel_acc_unet:>10.1%} {np.mean(ious_unet):>10.3f}')
                    pixel acc   mean IoU
     FCN (no skip)      98.1%      0.870
      U-Net (skip)      99.5%      0.943
In [6]:
fig, axes = plt.subplots(4, 4, figsize=(9, 9))
for i in range(4):
    axes[0, i].imshow(Xte[i], cmap='gray')
    axes[1, i].imshow(Mte[i], cmap='viridis', vmin=0, vmax=2)
    axes[2, i].imshow(preds_fcn[i], cmap='viridis', vmin=0, vmax=2)
    axes[3, i].imshow(preds_unet[i], cmap='viridis', vmin=0, vmax=2)
    for r in range(4):
        axes[r, i].axis('off')
for r, name in enumerate(['input', 'ground truth mask', 'pred (FCN)', 'pred (U-Net)']):
    axes[r, 0].set_title(name, fontsize=9, loc='left')
plt.tight_layout()
plt.show()
No description has been provided for this image

For small, tightly packed shapes, skip connections provide a substantial mean-IoU improvement over a plain FCN. An FCN cannot recover detail that was lost in the low-resolution bottleneck, because interpolation can only work on existing information. Skip connections fix this problem by giving the decoder a second, non-bottlenecked path to the encoder's high-resolution features, so precise boundaries never have to survive the bottleneck in the first place.

This helps explain why U-Net remains popular for tasks requiring precise boundaries, especially medical imaging. But skip connections are not universal: popular networks like plain FCN and DeepLabV3 have none, and even DeepLabV3+ adds only a single one, not a full U-Net-style connection at every resolution. So while skip connections are a useful tool, they are not a requirement — their value depends on how much fine spatial detail the task demands.

Exercises

  1. Increase n_shapes from 5 to 10, making the scene more crowded. Does the FCNTiny-vs-UNetTiny gap in mean IoU get larger or smaller? What does that suggest about when skip connections matter most?
  2. This lesson's loss is plain per-pixel cross-entropy. Print the per-pixel class weights implied by class pixel fractions above and try F.cross_entropy(logits, Mt, weight=inverse_class_freq) (Lesson 36's imbalance fix, applied here) to see whether it changes the circle/square IoU balance.
  3. nn.Upsample(mode='nearest') was used for simplicity. Try mode='bilinear', align_corners=False instead (Lesson 9's bilinear interpolation, now inside a network) and compare mean IoU for both FCNTiny and UNetTiny. Does the smoother upsampling help more or less than the skip connection does?
  4. Print the confidence (softmax probability of the predicted class) at a few pixels along the seam where the real FCN result switches from train to bus on the tram. Is the model confidently split, or genuinely uncertain right at that boundary — and what would you expect from a network being asked to choose between two classes when the true object matches neither one exactly?
Take the Lesson 42 Quiz →