Semantic Segmentation with U-Net in PyTorch: Implementation Guide for AI for Beginners
U-Net in PyTorch for semantic segmentation uses an encoder-decoder architecture with skip connections, implemented in the AI-For-Beginners repository using bilinear upsampling and batch normalization to segment PH² dermoscopy images at 256×256 resolution.
The Microsoft AI-For-Beginners curriculum provides a hands-on implementation of semantic segmentation using the classic U-Net architecture. Located in lessons/4-ComputerVision/12-Segmentation/SemanticSegmentationPytorch.ipynb, this PyTorch example demonstrates how to perform pixel-wise classification on medical images using an encoder-decoder CNN with skip connections.
U-Net Architecture Components
The implementation follows the original U-Net paper with a contracting encoder path, a bottleneck, and an expanding decoder path that restores spatial resolution through skip connections.
Encoder Path (Contracting)
The encoder consists of four convolutional blocks that progressively halve spatial dimensions while increasing channel depth. Each block in SemanticSegmentationPytorch.ipynb applies Conv2d → ReLU → BatchNorm2d → MaxPool2d(2):
enc_conv0: 3 input channels → 16 output channelsenc_conv1: 16 → 32 channelsenc_conv2: 32 → 64 channelsenc_conv3: 64 → 128 channels
This downsampling extracts multiscale features from the input dermoscopy images, reducing the 256×256 input to 16×16 feature maps at the deepest encoder level.
Bottleneck
Between the encoder and decoder, a bottleneck convolution expands the representation to 256 channels:
self.bottleneck_conv = nn.Conv2d(128, 256, kernel_size=3, padding=1)
This layer captures the highest-level semantic information before the decoder begins upsampling.
Decoder Path (Expanding)
The decoder mirrors the encoder structure using UpsamplingBilinear2d layers to double spatial resolution at each step. Critical to semantic segmentation performance, each decoder block concatenates the upsampled features with the corresponding encoder output (skip connections) before applying convolution:
dec_conv0: Concatenates withenc_conv3output (384 → 128 channels)dec_conv1: Concatenates withenc_conv2output (192 → 64 channels)dec_conv2: Concatenates withenc_conv1output (96 → 32 channels)dec_conv3: Concatenates withenc_conv0output (48 → 1 channel) with Sigmoid activation
The skip connections preserve spatial detail lost during pooling, enabling precise boundary delineation in the final segmentation mask.
Forward Pass Implementation
According to the source code in SemanticSegmentationPytorch.ipynb, the forward() method explicitly saves intermediate encoder outputs for skip connections before pooling operations:
def forward(self, x):
# Encoder with pooling
e0 = self.pool0(self.bn0(self.act0(self.enc_conv0(x))))
e1 = self.pool1(self.bn1(self.act1(self.enc_conv1(e0))))
e2 = self.pool2(self.bn2(self.act2(self.enc_conv2(e1))))
e3 = self.pool3(self.bn3(self.act3(self.enc_conv3(e2))))
# Save raw encoder outputs for skip connections
cat0 = self.bn0(self.act0(self.enc_conv0(x)))
cat1 = self.bn1(self.act1(self.enc_conv1(e0)))
cat2 = self.bn2(self.act2(self.enc_conv2(e1)))
cat3 = self.bn3(self.act3(self.enc_conv3(e2)))
# Bottleneck
b = self.bottleneck_conv(e3)
# Decoder with skip connections
d0 = self.dec_bn0(self.dec_conv0(
torch.cat((self.upsample0(b), cat3), dim=1)))
d1 = self.dec_bn1(self.dec_conv1(
torch.cat((self.upsample1(d0), cat2), dim=1)))
d2 = self.dec_bn2(self.dec_conv2(
torch.cat((self.upsample2(d1), cat1), dim=1)))
d3 = self.sigmoid(self.dec_conv3(
torch.cat((self.upsample3(d2), cat0), dim=1)))
return d3
The final output tensor has shape (batch, 1, 256, 256) containing per-pixel probabilities for the foreground class.
Training Configuration and Loss Function
The notebook configures training for binary semantic segmentation using standard PyTorch optimizers and loss functions:
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = UNet().to(device)
optimizer = optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-6)
criterion = nn.BCEWithLogitsLoss()
Adam optimizer settings include a learning rate of 0.001 and weight decay of 1e-6 for regularization. The training loop runs for 30 epochs, logging training and validation loss per epoch. For inference, the model applies a 0.5 threshold to the sigmoid output to generate binary masks:
predictions = (model(img).detach().cpu()[0] > 0.5).float()
Complete U-Net Code Example
Below is the full U-Net implementation extracted from SemanticSegmentationPytorch.ipynb, ready for adaptation to custom datasets:
import torch
import torch.nn as nn
import torch.optim as optim
class UNet(nn.Module):
def __init__(self):
super().__init__()
# Encoder blocks
self.enc_conv0 = nn.Conv2d(3, 16, 3, padding=1)
self.act0 = nn.ReLU()
self.bn0 = nn.BatchNorm2d(16)
self.pool0 = nn.MaxPool2d(2)
self.enc_conv1 = nn.Conv2d(16, 32, 3, padding=1)
self.act1 = nn.ReLU()
self.bn1 = nn.BatchNorm2d(32)
self.pool1 = nn.MaxPool2d(2)
self.enc_conv2 = nn.Conv2d(32, 64, 3, padding=1)
self.act2 = nn.ReLU()
self.bn2 = nn.BatchNorm2d(64)
self.pool2 = nn.MaxPool2d(2)
self.enc_conv3 = nn.Conv2d(64, 128, 3, padding=1)
self.act3 = nn.ReLU()
self.bn3 = nn.BatchNorm2d(128)
self.pool3 = nn.MaxPool2d(2)
# Bottleneck
self.bottleneck_conv = nn.Conv2d(128, 256, 3, padding=1)
# Decoder blocks with upsampling
self.upsample0 = nn.UpsamplingBilinear2d(scale_factor=2)
self.dec_conv0 = nn.Conv2d(384, 128, 3, padding=1)
self.dec_bn0 = nn.BatchNorm2d(128)
self.upsample1 = nn.UpsamplingBilinear2d(scale_factor=2)
self.dec_conv1 = nn.Conv2d(192, 64, 3, padding=1)
self.dec_bn1 = nn.BatchNorm2d(64)
self.upsample2 = nn.UpsamplingBilinear2d(scale_factor=2)
self.dec_conv2 = nn.Conv2d(96, 32, 3, padding=1)
self.dec_bn2 = nn.BatchNorm2d(32)
self.upsample3 = nn.UpsamplingBilinear2d(scale_factor=2)
self.dec_conv3 = nn.Conv2d(48, 1, 1)
self.sigmoid = nn.Sigmoid()
def forward(self, x):
# Encoder
e0 = self.pool0(self.bn0(self.act0(self.enc_conv0(x))))
e1 = self.pool1(self.bn1(self.act1(self.enc_conv1(e0))))
e2 = self.pool2(self.bn2(self.act2(self.enc_conv2(e1))))
e3 = self.pool3(self.bn3(self.act3(self.enc_conv3(e2))))
# Skip connection features
cat0 = self.bn0(self.act0(self.enc_conv0(x)))
cat1 = self.bn1(self.act1(self.enc_conv1(e0)))
cat2 = self.bn2(self.act2(self.enc_conv2(e1)))
cat3 = self.bn3(self.act3(self.enc_conv3(e2)))
# Bottleneck
b = self.bottleneck_conv(e3)
# Decoder with concatenation
d0 = self.dec_bn0(self.dec_conv0(
torch.cat((self.upsample0(b), cat3), dim=1)))
d1 = self.dec_bn1(self.dec_conv1(
torch.cat((self.upsample1(d0), cat2), dim=1)))
d2 = self.dec_bn2(self.dec_conv2(
torch.cat((self.upsample2(d1), cat1), dim=1)))
d3 = self.sigmoid(self.dec_conv3(
torch.cat((self.upsample3(d2), cat0), dim=1)))
return d3
Summary
- U-Net architecture in AI-For-Beginners uses a 4-layer encoder-decoder with skip connections to preserve spatial information for accurate semantic segmentation.
- The implementation processes 256×256 PH² dermoscopy images, outputting single-channel probability maps via Sigmoid activation.
- Bilinear upsampling doubles resolution at each decoder step, while BatchNorm2d and ReLU activations stabilize training throughout the network.
- Training uses BCEWithLogitsLoss with the Adam optimizer, achieving approximately 0.57 training loss after 30 epochs in the notebook example.
- Source files include
SemanticSegmentationPytorch.ipynbfor PyTorch andSemanticSegmentationTF.ipynbfor TensorFlow comparison.
Frequently Asked Questions
What dataset does the AI for Beginners U-Net use?
The implementation uses the PH² dataset, a collection of dermoscopy images for skin lesion analysis. The notebook loads these images alongside binary segmentation masks, resizing all inputs to 256×256 pixels to match the network's input requirements.
Why does the U-Net use skip connections in semantic segmentation?
Skip connections concatenate high-resolution encoder features with upsampled decoder features, allowing the network to recover fine-grained spatial details that are lost during max-pooling operations. This mechanism is essential for precise boundary detection in semantic segmentation tasks.
What loss function is used for semantic segmentation in this implementation?
The notebook uses Binary Cross-Entropy with Logits Loss (nn.BCEWithLogitsLoss()), which combines a sigmoid layer with binary cross-entropy in a numerically stable formulation. This is appropriate for the single-class segmentation task distinguishing lesions from background.
How does the decoder upsample feature maps in this U-Net?
The decoder utilizes bilinear interpolation via nn.UpsamplingBilinear2d(scale_factor=2) to double spatial dimensions at each of the four decoder stages. This approach is computationally efficient compared to transposed convolutions while maintaining segmentation quality for medical imaging applications.
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 →