用PyTorch从零实现FCN-8s语义分割:实战指南与VGG16迁移学习深度解析

语义分割作为计算机视觉领域的核心技术,正在自动驾驶、医疗影像分析等领域发挥越来越重要的作用。而FCN(全卷积网络)作为该领域的里程碑式模型,至今仍是理解现代分割算法的基础框架。本文将带您从零开始,用PyTorch实现一个完整的FCN-8s模型,并深入探讨如何利用VGG16预训练模型进行高效迁移学习。

1. 环境准备与核心概念

在开始编码之前,我们需要明确几个关键概念。FCN的核心创新在于将传统CNN中的全连接层替换为卷积层,使网络能够接受任意尺寸的输入并输出相同尺寸的分割图。这种"全卷积"的特性使其特别适合像素级预测任务。

必备工具安装

conda create -n fcn python=3.8
conda activate fcn
pip install torch torchvision matplotlib numpy

FCN系列模型(32s、16s、8s)的区别主要在于跳层连接(skip connection)的使用方式:

  • FCN-32s:仅使用最深层的特征图
  • FCN-16s:融合深层和中间层特征
  • FCN-8s:融合深层、中层和浅层特征

提示:使用Anaconda管理环境可以避免依赖冲突,特别是CUDA版本与PyTorch的兼容性问题

2. VGG16骨干网络改造

VGG16作为FCN的经典backbone,我们需要对其结构进行针对性改造。原始VGG16包含13个卷积层和3个全连接层,我们需要:

  1. 移除最后的全连接层
  2. 将全连接层替换为等效的1x1卷积
  3. 保留前面卷积层作为特征提取器
import torch
from torchvision import models

# 加载预训练VGG16(带BN版本)
pretrained_vgg = models.vgg16_bn(pretrained=True)

# 分解VGG16的各阶段特征提取层
class VGG16_FCN(nn.Module):
    def __init__(self):
        super().__init__()
        # 划分五个特征提取阶段
        self.stage1 = pretrained_vgg.features[:7]   # 到第一个池化层
        self.stage2 = pretrained_vgg.features[7:14] # 到第二个池化层
        self.stage3 = pretrained_vgg.features[14:24] # 到第三个池化层
        self.stage4 = pretrained_vgg.features[24:34] # 到第四个池化层
        self.stage5 = pretrained_vgg.features[34:]   # 到最后
        
    def forward(self, x):
        # 保留各阶段输出用于跳层连接
        s1 = self.stage1(x)
        s2 = self.stage2(s1)
        s3 = self.stage3(s2)
        s4 = self.stage4(s3)
        s5 = self.stage5(s4)
        return s3, s4, s5

关键改造点

  • 将原始VGG16的MaxPooling层保留作为下采样器
  • 分离不同深度的特征图用于后续融合
  • 冻结浅层网络权重(可选,根据数据集大小决定)

3. FCN-8s网络架构实现

FCN-8s的精髓在于多层次特征融合,我们需要实现三个关键组件:

  1. 1x1卷积:调整通道数匹配
  2. 转置卷积:实现特征图上采样
  3. 跳层连接:融合不同尺度特征
class FCN8s(nn.Module):
    def __init__(self, num_classes):
        super().__init__()
        self.vgg = VGG16_FCN()
        
        # 通道调整卷积
        self.conv1 = nn.Conv2d(512, 256, kernel_size=1)
        self.conv2 = nn.Conv2d(256, num_classes, kernel_size=1)
        
        # 转置卷积(上采样)
        self.upsample2x = nn.ConvTranspose2d(
            num_classes, num_classes, kernel_size=4, stride=2, padding=1)
        self.upsample8x = nn.ConvTranspose2d(
            num_classes, num_classes, kernel_size=16, stride=8, padding=4)
        
        # 初始化转置卷积核(双线性插值)
        self._init_upsample()
    
    def _init_upsample(self):
        """使用双线性插值初始化转置卷积"""
        kernel = self.bilinear_kernel(self.upsample2x.in_channels, 
                                     self.upsample2x.out_channels, 4)
        self.upsample2x.weight.data.copy_(kernel)
        
        kernel = self.bilinear_kernel(self.upsample8x.in_channels,
                                     self.upsample8x.out_channels, 16)
        self.upsample8x.weight.data.copy_(kernel)
    
    def bilinear_kernel(self, in_channels, out_channels, kernel_size):
        """生成双线性插值核"""
        factor = (kernel_size + 1) // 2
        center = factor - 0.5 if kernel_size % 2 == 0 else factor - 1
        
        og = np.ogrid[:kernel_size, :kernel_size]
        filt = (1 - abs(og[0] - center) / factor) * \
               (1 - abs(og[1] - center) / factor)
        
        weight = np.zeros((in_channels, out_channels, 
                          kernel_size, kernel_size), dtype='float32')
        weight[range(in_channels), range(out_channels), :, :] = filt
        return torch.from_numpy(weight)
    
    def forward(self, x):
        # 获取各阶段特征图
        s3, s4, s5 = self.vgg(x)
        
        # 第一级上采样:s5 2x -> 与s4融合
        s5_up = self.upsample2x(s5)
        add1 = s5_up + s4
        
        # 第二级上采样:融合结果 2x -> 与s3融合
        add1 = self.conv1(add1)
        add1_up = self.upsample2x(add1)
        add2 = add1_up + s3
        
        # 最终分类和上采样
        output = self.conv2(add2)
        output = self.upsample8x(output)
        return output

架构亮点

  • 使用双线性插值初始化转置卷积,加速训练收敛
  • 精确控制特征图尺寸匹配,确保跳层连接正确执行
  • 模块化设计便于调试和扩展

4. 训练策略与技巧

训练语义分割网络需要特别注意数据准备和超参数设置。以下是关键训练配置:

数据增强配置

from torchvision import transforms

train_transform = transforms.Compose([
    transforms.RandomHorizontalFlip(p=0.5),
    transforms.RandomVerticalFlip(p=0.5),
    transforms.RandomRotation(15),
    transforms.ColorJitter(
        brightness=0.1, contrast=0.1, saturation=0.1, hue=0.1),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], 
                        std=[0.229, 0.224, 0.225])
])

损失函数选择

criterion = nn.CrossEntropyLoss(
    weight=torch.tensor([1.0, 2.0, 1.5]),  # 类别权重
    ignore_index=255  # 忽略特定标签
)

优化器配置

optimizer = torch.optim.AdamW([
    {'params': model.vgg.parameters(), 'lr': 1e-4},  # 骨干网络较小学习率
    {'params': model.conv1.parameters()},
    {'params': model.conv2.parameters()},
    {'params': model.upsample2x.parameters()},
    {'params': model.upsample8x.parameters()}
], lr=1e-3, weight_decay=1e-4)

scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
    optimizer, mode='max', factor=0.5, patience=3)

训练循环关键代码

for epoch in range(epochs):
    model.train()
    for images, masks in train_loader:
        optimizer.zero_grad()
        outputs = model(images.cuda())
        loss = criterion(outputs, masks.cuda())
        loss.backward()
        optimizer.step()
    
    # 验证阶段
    model.eval()
    with torch.no_grad():
        for val_images, val_masks in val_loader:
            val_outputs = model(val_images.cuda())
            val_loss = criterion(val_outputs, val_masks.cuda())
            # 计算IoU等指标...
    
    scheduler.step(val_iou)  # 根据验证指标调整学习率

注意:在使用预训练模型时,建议先冻结骨干网络训练几轮,再解冻进行微调,这样能获得更好的收敛效果

5. 模型评估与性能优化

评估语义分割模型需要专门的指标,常用的包括:

指标名称计算公式说明
Pixel Accuracy正确像素数/总像素数最直观但可能不平衡
Mean IoU各类IoU的平均值最常用指标
Frequency Weighted IoU按类别频率加权考虑类别不平衡

常见性能瓶颈与解决方案

  1. 细节丢失严重

    • 增加浅层特征的权重
    • 尝试更密集的跳层连接(如FCN-4s)
    • 使用注意力机制增强重要特征
  2. 小物体分割效果差

    • 在损失函数中增加小物体权重
    • 使用多尺度训练策略
    • 尝试空洞卷积保留分辨率
  3. 训练不稳定

    • 检查梯度流动(特别是转置卷积部分)
    • 使用梯度裁剪
    • 尝试不同的学习率策略
# 多尺度推理示例
def multi_scale_inference(model, image, scales=[0.5, 1.0, 1.5]):
    outputs = []
    for scale in scales:
        resized_img = F.interpolate(
            image, scale_factor=scale, mode='bilinear', align_corners=False)
        outputs.append(F.interpolate(
            model(resized_img), 
            size=image.shape[2:], 
            mode='bilinear', 
            align_corners=False))
    return torch.mean(torch.stack(outputs), dim=0)

在实际项目中,我发现将转置卷积替换为最近邻上采样+普通卷积的组合有时能获得更锐利的边缘效果。此外,在数据增强中加入随机裁剪时,确保裁剪尺寸不小于原始图像的1/4,避免重要上下文信息丢失。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐