# Implementing Semantic Segmentation with U-Net Architecture: A Complete Guide

> Learn to implement semantic segmentation with U-Net architecture. This guide details encoder-decoder CNNs with skip connections for precise pixel-level classification, as shown in Microsoft's AI-For-Beginners.

- Repository: [Microsoft/AI-For-Beginners](https://github.com/microsoft/AI-For-Beginners)
- Tags: deep-dive
- Published: 2026-08-26

---

**Implementing semantic segmentation with U-Net architecture requires constructing an encoder-decoder convolutional neural network with skip connections that concatenate high-resolution encoder features with upsampled decoder outputs, enabling precise pixel-level classification as demonstrated in Microsoft's AI-For-Beginners repository.**

The Microsoft AI-For-Beginners curriculum provides a practical walkthrough of this computer vision task in `lessons/4-ComputerVision/12-Segmentation`, offering complete implementations in both TensorFlow and PyTorch. This guide examines the U-Net model structure, the specific code implementation, and the training pipeline used for binary image segmentation.

## Understanding the U-Net Architecture

The U-Net architecture follows an encoder-decoder design with symmetric skip connections that preserve spatial detail while learning hierarchical features. The implementation in AI-For-Beginners adheres to the original Ronneberger et al. (2015) specification with four encoder stages and corresponding decoder stages.

### Encoder (Contracting Path)

The encoder progressively down-samples the input through a series of **Conv2D → BatchNormalization → ReLU → MaxPool2D** blocks. Each stage doubles the number of filters while halving the spatial dimensions:

- **Stage 0**: 16 filters
- **Stage 1**: 32 filters  
- **Stage 2**: 64 filters
- **Stage 3**: 128 filters

This contraction extracts increasingly abstract feature representations from the input image.

### Bottleneck

A 3×3 convolution with **256 filters** bridges the encoder and decoder, capturing the most compressed high-level features before expansion begins.

### Decoder (Expanding Path)

The decoder uses **bilinear upsampling** followed by convolution blocks to restore spatial resolution. At each level, the upsampled features are concatenated with the corresponding encoder feature maps via skip connections, allowing the network to recover fine-grained details lost during pooling.

### Final Classification Layer

A 1×1 convolution reduces the channel dimension to **1** (for binary segmentation), producing a pixel-wise probability map matching the input resolution.

## TensorFlow Implementation

The complete TensorFlow implementation resides in `lessons/4-ComputerVision/12-Segmentation/SemanticSegmentationTF.ipynb`. The `UNet` class inherits from `tf.keras.Model` and explicitly defines layers in `__init__` with a custom `call` method for the forward pass.

### Model Definition

```python
import tensorflow as tf
from tensorflow import keras

class UNet(tf.keras.Model):
    def __init__(self):
        super().__init__()
        # Encoder blocks

        self.enc_conv0 = keras.Conv2D(16, kernel_size=3, padding='same')
        self.bn0       = keras.BatchNormalization()
        self.relu0    = keras.Activation('relu')
        self.pool0    = keras.MaxPool2D()

        self.enc_conv1 = keras.Conv2D(32, kernel_size=3, padding='same')
        self.bn1       = keras.BatchNormalization()
        self.relu1    = keras.Activation('relu')
        self.pool1    = keras.MaxPool2D()

        self.enc_conv2 = keras.Conv2D(64, kernel_size=3, padding='same')
        self.bn2       = keras.BatchNormalization()
        self.relu2    = keras.Activation('relu')
        self.pool2    = keras.MaxPool2D()

        self.enc_conv3 = keras.Conv2D(128, kernel_size=3, padding='same')
        self.bn3       = keras.BatchNormalization()
        self.relu3    = keras.Activation('relu')
        self.pool3    = keras.MaxPool2D()

        # Bottleneck

        self.bottleneck_conv = keras.Conv2D(256, kernel_size=(3, 3), padding='same')

        # Decoder blocks (upsampling + conv)

        self.upsample0 = keras.UpSampling2D(interpolation='bilinear')
        self.dec_conv0 = keras.Conv2D(128, kernel_size=3, padding='same')
        self.dec_bn0   = keras.BatchNormalization()
        self.dec_relu0 = keras.Activation('relu')

        self.upsample1 = keras.UpSampling2D(interpolation='bilinear')
        self.dec_conv1 = keras.Conv2D(64, kernel_size=3, padding='same')
        self.dec_bn1   = keras.BatchNormalization()
        self.dec_relu1 = keras.Activation('relu')

        self.upsample2 = keras.UpSampling2D(interpolation='bilinear')
        self.dec_conv2 = keras.Conv2D(32, kernel_size=3, padding='same')
        self.dec_bn2   = keras.BatchNormalization()
        self.dec_relu2 = keras.Activation('relu')

        self.upsample3 = keras.UpSampling2D(interpolation='bilinear')
        self.dec_conv3 = keras.Conv2D(1, kernel_size=1)

        # Concatenation layers for skip connections

        self.cat0 = keras.Concatenate(axis=3)
        self.cat1 = keras.Concatenate(axis=3)
        self.cat2 = keras.Concatenate(axis=3)
        self.cat3 = keras.Concatenate(axis=3)

    def call(self, input):
        # Encoder forward pass

        e0 = self.pool0(self.relu0(self.bn0(self.enc_conv0(input))))
        e1 = self.pool1(self.relu1(self.bn1(self.enc_conv1(e0))))
        e2 = self.pool2(self.relu2(self.bn2(self.enc_conv2(e1))))
        e3 = self.pool3(self.relu3(self.bn3(self.enc_conv3(e2))))

        # Preserve low‑level features for skip connections

        cat0 = self.relu0(self.bn0(self.enc_conv0(input)))
        cat1 = self.relu1(self.bn1(self.enc_conv1(e0)))
        cat2 = self.relu2(self.bn2(self.enc_conv2(e1)))
        cat3 = self.relu3(self.bn3(self.enc_conv3(e2)))

        # Bottleneck

        b = self.bottleneck_conv(e3)

        # Decoder with skip connections

        cat_tens0 = self.cat0([self.upsample0(b), cat3])
        d0 = self.dec_relu0(self.dec_bn0(self.dec_conv0(cat_tens0)))

        cat_tens1 = self.cat1([self.upsample1(d0), cat2])
        d1 = self.dec_relu1(self.dec_bn1(self.dec_conv1(cat_tens1)))

        cat_tens2 = self.cat2([self.upsample2(d1), cat1])
        d2 = self.dec_relu2(self.dec_bn2(self.dec_conv2(cat_tens2)))

        cat_tens3 = self.cat3([self.upsample3(d2), cat0])
        d3 = self.dec_conv3(cat_tens3)

        return d3

```

### Compilation and Training

The training pipeline uses the PH2 skin-lesion dataset with images resized to **256×256** pixels and binarized masks. The model employs the Adam optimizer with specific hyperparameters tuned for this segmentation task:

```python
from tensorflow.keras import optimizers, losses

# Instantiate the model

model = UNet()

# Optimizer with weight decay

optimizer = optimizers.Adam(learning_rate=3e-4, decay=8e-9)

# Binary crossentropy for pixel-wise classification

loss_fn = losses.BinaryCrossentropy(from_logits=True)

# Compile

model.compile(loss=loss_fn, optimizer=optimizer)

# Train for 100 epochs

model.fit(
    X_train, y_train,
    epochs=100,
    batch_size=64,
    validation_data=(X_test, y_test),
    shuffle=True
)

```

## PyTorch Implementation

An equivalent PyTorch implementation is available in `lessons/4-ComputerVision/12-Segmentation/SemanticSegmentationPytorch.ipynb`. The architecture mirrors the TensorFlow version using `nn.Conv2d`, `nn.BatchNorm2d`, and `nn.Upsample` with bilinear interpolation:

```python
import torch
import torch.nn as nn

class UNet(nn.Module):
    def __init__(self):
        super().__init__()
        # Encoder blocks

        self.enc_conv0 = nn.Conv2d(3, 16, 3, padding=1)
        self.bn0       = nn.BatchNorm2d(16)
        self.relu0    = nn.ReLU()
        self.pool0    = nn.MaxPool2d(2)

        # ... (additional encoder layers follow same pattern) ...

        # Bottleneck

        self.bottleneck = nn.Conv2d(128, 256, 3, padding=1)

        # Decoder blocks

        self.up0   = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)
        self.dec0  = nn.Conv2d(256 + 128, 128, 3, padding=1)
        self.bn0d  = nn.BatchNorm2d(128)
        self.relu0d = nn.ReLU()

        # Final 1x1 conv

        self.final = nn.Conv2d(16, 1, 1)

    def forward(self, x):
        # Encoder forward pass

        e0 = self.pool0(self.relu0(self.bn0(self.enc_conv0(x))))
        # ... (intermediate layers) ...

        
        # Decoder with concatenated skips

        d0 = self.relu0d(self.bn0d(self.dec0(torch.cat([self.up0(b), e3], dim=1))))
        # ... (remaining decoder layers) ...

        
        return self.final(d3)

```

## Dataset and Preprocessing

The implementation uses the **PH2 dataset** for skin lesion segmentation. The `load_dataset()` function handles:

1. Downloading and extracting the dataset
2. Resizing images to **256×256** resolution
3. Binarizing segmentation masks to create pixel-wise ground truth
4. Returning TensorFlow tensors or PyTorch tensors ready for training

This preprocessing ensures consistent input dimensions and normalized intensity values required by the U-Net architecture.

## Summary

- **U-Net architecture** combines an encoder-decoder structure with skip connections to preserve spatial precision while capturing semantic context.
- **TensorFlow implementation** in `SemanticSegmentationTF.ipynb` defines the model as a custom `tf.keras.Model` with explicit layer definitions and bilinear upsampling in the decoder.
- **Training configuration** uses Adam optimizer (`learning_rate=3e-4`, `decay=8e-9`), binary crossentropy loss, 100 epochs, and batch size 64 on 256×256 images.
- **PyTorch alternative** in `SemanticSegmentationPytorch.ipynb` provides identical functionality using `nn.Module` and `torch.cat` for skip connections.
- **Skip connections** are implemented via `keras.Concatenate` (TensorFlow) or `torch.cat` (PyTorch) at each decoder stage, concatenating encoder features with upsampled decoder outputs.

## Frequently Asked Questions

### What is the purpose of skip connections in U-Net?

Skip connections preserve high-resolution spatial information from the encoder that would otherwise be lost during max-pooling operations. By concatenating these features with the corresponding decoder layers, the network can accurately localize object boundaries while maintaining the semantic understanding learned in deeper layers.

### Why does the implementation use bilinear upsampling instead of transposed convolutions?

Bilinear upsampling followed by standard convolution avoids **checkerboard artifacts** that commonly occur with transposed convolutions (deconvolutions). The `interpolation='bilinear'` parameter in `keras.UpSampling2D` provides smoother feature map expansion, leading to more stable training and cleaner segmentation masks.

### How is the PH2 dataset prepared for the segmentation task?

The dataset preprocessing in `load_dataset()` resizes all images to **256×256** pixels and normalizes pixel values. Segmentation masks are binarized to create binary classification targets where each pixel is classified as either lesion or background, suitable for the single-channel output of the U-Net model.

### Can this U-Net implementation handle multi-class segmentation?

Yes, the architecture supports multi-class segmentation by modifying the final convolution layer to output channels equal to the number of classes (instead of 1) and changing the loss function to categorical crossentropy. However, the current AI-For-Beginners implementation specifically demonstrates binary segmentation for skin lesion detection.