{ "cells": [ { "cell_type": "markdown", "id": "95000fa8", "metadata": {}, "source": [ "# Lesson 42: Semantic Segmentation\n", "\n", "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." ] }, { "cell_type": "code", "execution_count": null, "id": "b259f7a5", "metadata": {}, "outputs": [], "source": [ "import numpy as np\n", "import torch\n", "import torch.nn as nn\n", "import torch.nn.functional as F\n", "import matplotlib.pyplot as plt" ] }, { "cell_type": "markdown", "id": "fbcfd6d4", "metadata": {}, "source": [ "## A 3-class pixel labeling task\n", "\n", "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." ] }, { "cell_type": "code", "execution_count": null, "id": "4042a733", "metadata": {}, "outputs": [], "source": [ "def make_scene(rng, size=32, n_shapes=5):\n", " scene = np.zeros((size, size), dtype=np.float32)\n", " mask = np.zeros((size, size), dtype=np.int64) # 0=background, 1=circle, 2=square\n", " for _ in range(n_shapes):\n", " shape_type = rng.choice([1, 2])\n", " r = rng.integers(2, 4)\n", " cx, cy = rng.integers(r, size - r), rng.integers(r, size - r)\n", " yy, xx = np.mgrid[0:size, 0:size]\n", " if shape_type == 1:\n", " m = ((xx - cx) ** 2 + (yy - cy) ** 2) <= r ** 2\n", " else:\n", " m = (np.abs(xx - cx) <= r) & (np.abs(yy - cy) <= r)\n", " scene[m] = 1.0\n", " mask[m] = shape_type\n", " scene = np.clip(scene + rng.normal(0, 0.05, scene.shape), 0, 1).astype(np.float32)\n", " return scene, mask\n", "\n", "rng = np.random.default_rng(13)\n", "N = 300\n", "scenes, masks = [], []\n", "for _ in range(N):\n", " s, m = make_scene(rng)\n", " scenes.append(s); masks.append(m)\n", "scenes = np.array(scenes, dtype=np.float32)\n", "masks = np.array(masks, dtype=np.int64)\n", "\n", "split = int(0.85 * N)\n", "Xtr, Mtr = scenes[:split], masks[:split]\n", "Xte, Mte = scenes[split:], masks[split:]\n", "\n", "class_fracs = (np.bincount(Mtr.ravel()) / Mtr.size).round(3).tolist()\n", "print('class pixel fractions (train):', dict(zip(['background', 'circle', 'square'], class_fracs)))\n", "\n", "fig, axes = plt.subplots(2, 4, figsize=(9, 4.5))\n", "for i in range(4):\n", " axes[0, i].imshow(Xtr[i], cmap='gray'); axes[0, i].axis('off')\n", " axes[1, i].imshow(Mtr[i], cmap='viridis', vmin=0, vmax=2); axes[1, i].axis('off')\n", "axes[0, 0].set_title('input', fontsize=9, loc='left')\n", "axes[1, 0].set_title('per-pixel ground truth label', fontsize=9, loc='left')\n", "plt.show()" ] }, { "cell_type": "markdown", "id": "0c655321", "metadata": {}, "source": [ "## A minimal FCN-style network\n", "\n", "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." ] }, { "cell_type": "code", "execution_count": null, "id": "aab2f1d6", "metadata": {}, "outputs": [], "source": [ "class FCNTiny(nn.Module):\n", " def __init__(self, n_classes=3):\n", " super().__init__()\n", " self.enc1 = nn.Sequential(nn.Conv2d(1, 16, 3, padding=1), nn.ReLU())\n", " self.enc2 = nn.Sequential(nn.Conv2d(16, 32, 3, padding=1), nn.ReLU())\n", " self.enc3 = nn.Sequential(nn.Conv2d(32, 64, 3, padding=1), nn.ReLU())\n", " self.pool = nn.MaxPool2d(2)\n", " self.up = nn.Upsample(scale_factor=2, mode='nearest')\n", " self.dec2 = nn.Sequential(nn.Conv2d(64, 32, 3, padding=1), nn.ReLU())\n", " self.dec1 = nn.Sequential(nn.Conv2d(32, 16, 3, padding=1), nn.ReLU())\n", " self.out = nn.Conv2d(16, n_classes, 1)\n", "\n", " def forward(self, x):\n", " f1 = self.enc1(x) # (B,16,H,W)\n", " f2 = self.enc2(self.pool(f1)) # (B,32,H/2,W/2)\n", " f3 = self.enc3(self.pool(f2)) # (B,64,H/4,W/4) -- the bottleneck\n", " d2 = self.dec2(self.up(f3)) # upsample only -- no fusion with f2\n", " d1 = self.dec1(self.up(d2)) # upsample only -- no fusion with f1\n", " return self.out(d1)\n", "\n", "def mean_iou(preds, targets, n_classes=3):\n", " ious = []\n", " for c in range(n_classes):\n", " p, t = preds == c, targets == c\n", " inter, union = (p & t).sum(), (p | t).sum()\n", " ious.append(inter / union if union > 0 else float('nan'))\n", " return ious\n", "\n", "torch.manual_seed(0)\n", "model_fcn = FCNTiny()\n", "opt = torch.optim.Adam(model_fcn.parameters(), lr=0.01)\n", "Xt = torch.tensor(Xtr).unsqueeze(1); Mt = torch.tensor(Mtr)\n", "for _ in range(200):\n", " opt.zero_grad()\n", " loss = F.cross_entropy(model_fcn(Xt), Mt)\n", " loss.backward()\n", " opt.step()\n", "\n", "with torch.no_grad():\n", " preds_fcn = model_fcn(torch.tensor(Xte).unsqueeze(1)).argmax(1).numpy()\n", "\n", "pixel_acc_fcn = (preds_fcn == Mte).mean()\n", "ious_fcn = mean_iou(preds_fcn, Mte)\n", "print(f'pixel accuracy: {pixel_acc_fcn:.1%}')\n", "for name, iou in zip(['background', 'circle', 'square'], ious_fcn):\n", " print(f' {name:>10} IoU: {iou:.3f}')\n", "print(f'mean IoU: {np.mean(ious_fcn):.3f}')" ] }, { "cell_type": "markdown", "id": "3f254359", "metadata": {}, "source": [ "## In practice: FCN on a real photo\n", "\n", "`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." ] }, { "cell_type": "code", "execution_count": null, "id": "e0968f98", "metadata": {}, "outputs": [], "source": [ "import torchvision\n", "import cv2\n", "import matplotlib.patches as mpatches\n", "\n", "weights = torchvision.models.segmentation.FCN_ResNet50_Weights.COCO_WITH_VOC_LABELS_V1\n", "fcn = torchvision.models.segmentation.fcn_resnet50(weights=weights)\n", "fcn.eval() # frozen, no training at all\n", "voc_classes = weights.meta['categories']\n", "\n", "photo = cv2.imread('../img/bicycle_train.jpg')\n", "photo_rgb = cv2.cvtColor(photo, cv2.COLOR_BGR2RGB)\n", "\n", "preprocess = weights.transforms()\n", "photo_tensor = torch.tensor(photo_rgb / 255.0, dtype=torch.float32).permute(2, 0, 1)\n", "batch = preprocess(photo_tensor).unsqueeze(0)\n", "\n", "with torch.no_grad():\n", " real_out = fcn(batch)['out']\n", "real_pred = real_out.argmax(1)[0].numpy()\n", "# the model's internal resolution differs from the photo's; resize the label map back up to match\n", "real_pred_full = cv2.resize(real_pred.astype(np.uint8), (photo_rgb.shape[1], photo_rgb.shape[0]),\n", " interpolation=cv2.INTER_NEAREST)\n", "\n", "found_ids = sorted(c for c in np.unique(real_pred_full) if c != 0)\n", "colors = plt.cm.tab10(np.linspace(0, 1, 10))\n", "\n", "label_rgb = np.full((*real_pred_full.shape, 3), 255, dtype=np.uint8) # background = white\n", "for i, cls_id in enumerate(found_ids):\n", " color_rgb = (np.array(colors[i][:3]) * 255).astype(np.uint8)\n", " label_rgb[real_pred_full == cls_id] = color_rgb\n", "\n", "fig, axes = plt.subplots(1, 2, figsize=(12, 4.5))\n", "axes[0].imshow(photo_rgb); axes[0].set_title('input'); axes[0].axis('off')\n", "axes[1].imshow(label_rgb); axes[1].set_title('FCN per-pixel labels'); axes[1].axis('off')\n", "legend_handles = [mpatches.Patch(color=colors[i], label=voc_classes[c]) for i, c in enumerate(found_ids)]\n", "axes[1].legend(handles=legend_handles, loc='upper right', fontsize=8)\n", "plt.tight_layout()\n", "plt.show()" ] }, { "cell_type": "markdown", "id": "2900630a", "metadata": {}, "source": [ "
Photo by Felix Hanspach on Unsplash
" ] }, { "cell_type": "markdown", "id": "ae9b4441", "metadata": {}, "source": [ "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." ] }, { "cell_type": "markdown", "id": "4433b6cd", "metadata": {}, "source": [ "## Adding skip connections: U-Net\n", "\n", "`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." ] }, { "cell_type": "code", "execution_count": null, "id": "eb504d2b", "metadata": {}, "outputs": [], "source": [ "class UNetTiny(nn.Module):\n", " def __init__(self, n_classes=3):\n", " super().__init__()\n", " self.enc1 = nn.Sequential(nn.Conv2d(1, 16, 3, padding=1), nn.ReLU())\n", " self.enc2 = nn.Sequential(nn.Conv2d(16, 32, 3, padding=1), nn.ReLU())\n", " self.enc3 = nn.Sequential(nn.Conv2d(32, 64, 3, padding=1), nn.ReLU())\n", " self.pool = nn.MaxPool2d(2)\n", " self.up = nn.Upsample(scale_factor=2, mode='nearest')\n", " self.dec2 = nn.Sequential(nn.Conv2d(64 + 32, 32, 3, padding=1), nn.ReLU())\n", " self.dec1 = nn.Sequential(nn.Conv2d(32 + 16, 16, 3, padding=1), nn.ReLU())\n", " self.out = nn.Conv2d(16, n_classes, 1)\n", "\n", " def forward(self, x):\n", " f1 = self.enc1(x) # (B,16,H,W)\n", " f2 = self.enc2(self.pool(f1)) # (B,32,H/2,W/2)\n", " f3 = self.enc3(self.pool(f2)) # (B,64,H/4,W/4)\n", " d2 = self.dec2(torch.cat([self.up(f3), f2], dim=1)) # skip from f2\n", " d1 = self.dec1(torch.cat([self.up(d2), f1], dim=1)) # skip from f1\n", " return self.out(d1)\n", "\n", "torch.manual_seed(0)\n", "model_unet = UNetTiny()\n", "opt2 = torch.optim.Adam(model_unet.parameters(), lr=0.01)\n", "for _ in range(200):\n", " opt2.zero_grad()\n", " loss = F.cross_entropy(model_unet(Xt), Mt)\n", " loss.backward()\n", " opt2.step()\n", "\n", "with torch.no_grad():\n", " preds_unet = model_unet(torch.tensor(Xte).unsqueeze(1)).argmax(1).numpy()\n", "\n", "pixel_acc_unet = (preds_unet == Mte).mean()\n", "ious_unet = mean_iou(preds_unet, Mte)\n", "\n", "print(f'{\"\":>18} {\"pixel acc\":>10} {\"mean IoU\":>10}')\n", "print(f'{\"FCN (no skip)\":>18} {pixel_acc_fcn:>10.1%} {np.mean(ious_fcn):>10.3f}')\n", "print(f'{\"U-Net (skip)\":>18} {pixel_acc_unet:>10.1%} {np.mean(ious_unet):>10.3f}')" ] }, { "cell_type": "code", "execution_count": null, "id": "2fcb0fbc", "metadata": {}, "outputs": [], "source": [ "fig, axes = plt.subplots(4, 4, figsize=(9, 9))\n", "for i in range(4):\n", " axes[0, i].imshow(Xte[i], cmap='gray')\n", " axes[1, i].imshow(Mte[i], cmap='viridis', vmin=0, vmax=2)\n", " axes[2, i].imshow(preds_fcn[i], cmap='viridis', vmin=0, vmax=2)\n", " axes[3, i].imshow(preds_unet[i], cmap='viridis', vmin=0, vmax=2)\n", " for r in range(4):\n", " axes[r, i].axis('off')\n", "for r, name in enumerate(['input', 'ground truth mask', 'pred (FCN)', 'pred (U-Net)']):\n", " axes[r, 0].set_title(name, fontsize=9, loc='left')\n", "plt.tight_layout()\n", "plt.show()" ] }, { "cell_type": "markdown", "id": "b4cd6aaa", "metadata": {}, "source": [ "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.\n", "\n", "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." ] }, { "cell_type": "markdown", "id": "d58f6e24", "metadata": {}, "source": [ "### Exercises\n", "\n", "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?\n", "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.\n", "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?\n", "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?" ] } ], "metadata": { "kernelspec": { "display_name": "Python 3", "language": "python", "name": "python3" }, "language_info": { "name": "python", "version": "3.10.0" } }, "nbformat": 4, "nbformat_minor": 5 }