nitinvig's picture
Upload 6 files
dc9e606 verified
Raw
History Blame Contribute Delete
7.93 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
class DoubleConv(nn.Module):
"""
DoubleConv Module
=================
A standard building block for UNet, consisting of two consecutive convolution layers.
Each 3x3 convolution is followed by Batch Normalization and ReLU activation.
Structure:
Input -> [Conv3x3 -> BatchNorm -> ReLU] -> [Conv3x3 -> BatchNorm -> ReLU] -> Output
"""
def __init__(self, in_channels, out_channels):
super().__init__()
self.double_conv = nn.Sequential(
nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),
nn.BatchNorm2d(out_channels),
nn.ReLU(inplace=True),
nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),
nn.BatchNorm2d(out_channels),
nn.ReLU(inplace=True)
)
def forward(self, x):
return self.double_conv(x)
class Down(nn.Module):
"""
Down Module
===========
Handles the downsampling step in the encoder part of the UNet.
It supports two modes of downsampling:
1. 'maxpool': Uses MaxPool2d(2) to halve the spatial dimensions.
2. 'strided': Uses a Strided Conv (kernel=3, stride=2) to halve dimensions while learning features.
After downsampling, a DoubleConv block processes the features.
"""
def __init__(self, in_channels, out_channels, mode='maxpool'):
super().__init__()
self.mode = mode
if mode == 'maxpool':
# Option 1: MaxPool downsampling (Standard UNet)
self.down_layer = nn.MaxPool2d(2)
self.conv = DoubleConv(in_channels, out_channels)
elif mode == 'strided':
# Option 2: Strided Convolution downsampling
# Replaces the pooling operation with a learnable strided convolution
self.down_layer = nn.Sequential(
nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=2, padding=1),
nn.BatchNorm2d(out_channels),
nn.ReLU(inplace=True)
)
# The strided conv handles the channel change (in -> out).
# The following DoubleConv refines these features (out -> out).
self.conv = DoubleConv(out_channels, out_channels)
else:
raise ValueError(f"Unknown downsample mode: {mode}")
def forward(self, x):
if self.mode == 'maxpool':
x = self.down_layer(x)
return self.conv(x)
else:
# Strided path
x = self.down_layer(x) # [B, OutCh, H/2, W/2]
return self.conv(x) # [B, OutCh, H/2, W/2]
class Up(nn.Module):
"""
Up Module
=========
Handles the upsampling step in the decoder part of the UNet.
It supports two modes of upsampling:
1. 'transpose': Uses ConvTranspose2d to learn how to upsample.
2. 'upsample': Uses bilinear interpolation (nn.Upsample).
Steps:
1. Upsample the input tensor (x1) from the previous lower layer.
2. Concatenate it with the corresponding feature map from the encoder (x2) (Skip Connection).
- Handles padding if dimensions don't match perfectly.
3. Process the combined features with a DoubleConv block.
"""
def __init__(self, in_channels, out_channels, mode='transpose'):
super().__init__()
if mode == 'transpose':
# Option 1: Transpose Convolution
# Typical for original UNet. Upsamples and reduces channels by half.
# in_channels is the dimension of the deep feature map coming UP.
self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size=2, stride=2)
self.up_mode = 'transpose'
elif mode == 'upsample':
# Option 2: Bilinear Upsampling
# Does not reduce channels itself, so we need a 1x1 conv to reduce channels
# to match the skip connection size before DoubleConv.
self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)
self.conv_adjust = nn.Conv2d(in_channels, in_channels // 2, kernel_size=1)
self.up_mode = 'upsample'
else:
raise ValueError(f"Unknown upsample mode: {mode}")
# DoubleConv takes the concatenated input.
# Channels = (in_channels // 2 from Up) + (in_channels // 2 from Skip) = in_channels
# Outputs count = out_channels
self.conv = DoubleConv(in_channels, out_channels)
def forward(self, x1, x2):
"""
x1: Input from the previous decoder layer (to be upsampled)
x2: Input from the encoder layer (skip connection)
"""
x1 = self.up(x1)
if hasattr(self, 'conv_adjust'):
x1 = self.conv_adjust(x1)
# Handle padding if x1 and x2 have slightly different sizes due to odd dimensions
# input is CHW
diffY = x2.size()[2] - x1.size()[2]
diffX = x2.size()[3] - x1.size()[3]
x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2,
diffY // 2, diffY - diffY // 2])
# Concatenate x2 (skip) and x1 (upsampled) along the channel dimension
x = torch.cat([x2, x1], dim=1)
return self.conv(x)
class UNet(nn.Module):
"""
UNet Architecture
=================
A U-shaped encoder-decoder architecture for image segmentation.
Configurable Parameters:
- n_channels: Number of input image channels (e.g., 3 for RGB).
- n_classes: Number of output classes (e.g., 1 for binary mask).
- downsample_mode: 'maxpool' or 'strided'.
- upsample_mode: 'transpose' or 'upsample' (bilinear).
"""
def __init__(self, n_channels, n_classes, downsample_mode='maxpool', upsample_mode='transpose'):
super(UNet, self).__init__()
self.n_channels = n_channels
self.n_classes = n_classes
self.downsample_mode = downsample_mode
self.upsample_mode = upsample_mode
# Initial Feature Extraction
# Input: [B, n_channels, H, W] -> Output: [B, 64, H, W]
self.inc = DoubleConv(n_channels, 64)
# Encoder (Downsampling Path)
# Each step reduces H,W by 2 and doubles Channels
# Down 1: 64 -> 128
self.down1 = Down(64, 128, mode=downsample_mode)
# Down 2: 128 -> 256
self.down2 = Down(128, 256, mode=downsample_mode)
# Down 3: 256 -> 512
self.down3 = Down(256, 512, mode=downsample_mode)
# Bridge / Bottleneck
# Standard UNet goes to 1024.
self.down4 = Down(512, 1024, mode=downsample_mode)
# Decoder (Upsampling Path)
# Each step doubles H,W and halves Channels (logic handled in Up block)
self.up1 = Up(1024, 512, mode=upsample_mode)
self.up2 = Up(512, 256, mode=upsample_mode)
self.up3 = Up(256, 128, mode=upsample_mode)
self.up4 = Up(128, 64, mode=upsample_mode)
# Final Classification Layer
# Reduces 64 channels to n_classes (1 per pixel for binary)
self.outc = nn.Conv2d(64, n_classes, kernel_size=1)
def forward(self, x):
# Encoder Path with Skip Connections
x1 = self.inc(x) # [B, 64, H, W]
x2 = self.down1(x1) # [B, 128, H/2, W/2]
x3 = self.down2(x2) # [B, 256, H/4, W/4]
x4 = self.down3(x3) # [B, 512, H/8, W/8]
x5 = self.down4(x4) # [B, 1024, H/16, W/16] (Bottleneck)
# Decoder Path
# Pass skip connections (x4, x3, x2, x1) to Up modules
x = self.up1(x5, x4) # [B, 512, H/8, W/8]
x = self.up2(x, x3) # [B, 256, H/4, W/4]
x = self.up3(x, x2) # [B, 128, H/2, W/2]
x = self.up4(x, x1) # [B, 64, H, W]
logits = self.outc(x) # [B, n_classes, H, W]
return logits