# Semantic Segmentation with U-Net in PyTorch: Implementation Guide for AI for Beginners

> Learn semantic segmentation with U-Net in PyTorch. Implement this encoder-decoder model for image segmentation using the AI-For-Beginners repository. Perfect for beginners!

- Repository: [Microsoft/AI-For-Beginners](https://github.com/microsoft/AI-For-Beginners)
- Tags: how-to-guide
- Published: 2026-08-29

---

**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 channels
- `enc_conv1`: 16 → 32 channels  
- `enc_conv2`: 32 → 64 channels
- `enc_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:

```python
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 with `enc_conv3` output (384 → 128 channels)
- `dec_conv1`: Concatenates with `enc_conv2` output (192 → 64 channels)  
- `dec_conv2`: Concatenates with `enc_conv1` output (96 → 32 channels)
- `dec_conv3`: Concatenates with `enc_conv0` output (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:

```python
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:

```python
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:

```python
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:

```python
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.ipynb` for PyTorch and `SemanticSegmentationTF.ipynb` for 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.