Implementing Semantic Segmentation with U-Net Architecture: A Complete Guide
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
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:
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:
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:
- Downloading and extracting the dataset
- Resizing images to 256×256 resolution
- Binarizing segmentation masks to create pixel-wise ground truth
- 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.ipynbdefines the model as a customtf.keras.Modelwith 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.ipynbprovides identical functionality usingnn.Moduleandtorch.catfor skip connections. - Skip connections are implemented via
keras.Concatenate(TensorFlow) ortorch.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.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →