在深度学习驱动的计算机视觉领域,尤其是生物医学图像分割任务中,U-Net 无疑是一座里程碑式的神经网络架构。它凭借精巧的对称设计和高效的跳跃连接,解决了传统卷积网络在像素级预测时丢失空间细节的痛点。本文将带你从零开始,深入拆解 U-Net 的每一个组件,并通过 PyTorch 代码实践,彻底掌握这一经典模型的复现技巧。

为什么语义分割需要 U-Net?

经典的卷积神经网络(CNN)在设计上通常依赖逐层下采样来扩大感受野,提取高层语义特征。这种策略在图像分类任务中表现出色,但在需要输出逐像素标签的语义分割任务中,下采样过程会不可逆地丢失大量精细的空间位置信息。即便后续通过上采样操作恢复分辨率,也难以重新找回清晰的边界和局部结构,导致分割结果轮廓粗糙、定位不准。

U-Net 通过结构上的巧妙重组,在提取高层语义信息的同时,最大程度地保留并利用了高分辨率的空间特征。整体上,U-Net 采用了一种对称的 编码器-解码器 结构:左侧为逐步下采样的编码路径,用于提取多尺度语义特征;右侧为逐步上采样的解码路径,用于恢复空间分辨率并生成像素级预测结果。在网络最底部,两条路径通过一个瓶颈层相连。

在每一层中,网络将编码器尚未经过下采样的高分辨率特征,直接拼接到解码器对应层中。这种 跳跃连接 机制,使 U-Net 能够在保持强表达力的同时,实现对目标边界和细节结构的精准定位。这种结构设计思想,也为后续许多先进的 AI 模型提供了灵感。

语义分割与传统图像分类存在本质区别。图像分类关心这张图是什么,而语义分割要求对每一个像素进行判别,不仅需要理解目标是什么,还要准确回答它在什么位置边界在哪里。

编码器:多尺度语义特征的提取器

编码器的作用类似于一个特征金字塔,通过重复的池化和卷积操作,逐步将输入图像压缩为语义丰富但分辨率较低的特征图。其核心模块是 DoubleConvDown 下采样模块。

核心模块解析

  • DoubleConv:由两层 3×3 卷积组成,每层卷积后都接有批归一化(BN)和 ReLU 激活函数。由于设置了 padding=1,每次卷积都不会改变特征图的空间尺寸(H, W),仅改变通道数(C)。
  • Down:先通过一个步长为 2 的 2×2 最大池化将特征图尺寸减半,然后接一个 DoubleConv 模块,在新的尺度上继续提取特征。

以下方的典型配置为例,输入图像尺寸为 256×256×1:

  • 初始 DoubleConv:两次 3×3×64 卷积,输出 256×256×64。
  • down1:池化后尺寸减半,DoubleConv 将通道数提升至 128,输出 128×128×128。
  • down2:经过池化和 DoubleConv,输出 64×64×256。
  • down3:继续下采样,输出 32×32×512。
  • down4:瓶颈层,输出 16×16×512,作为解码器上采样的起点。

在下采样的同时,编码器每一级的输出特征都会被保留下来,供解码器通过跳跃连接使用。

编码器开始先做一次双卷积增加原始图像的通道数,随后重复 4 次下采样+双卷积下采样用 2×2 最大池化将分辨率减半;双卷积用两次 3×3 卷积(padding=1, stride=1)在该尺度上提取特征。具体参数如下:

在代码实现上,DoubleConv 模块的构建非常直观:

class DoubleConv(nn.Sequential):
    def __init__(self, in_channels, out_channels, mid_channels=None):
        if mid_channels is None:
            mid_channels = out_channels
        super(DoubleConv, self).__init__(
            nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=1, bias=False),
            nn.BatchNorm2d(mid_channels),
            nn.ReLU(inplace=True),
            nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1, bias=False),
            nn.BatchNorm2d(out_channels),
            nn.ReLU(inplace=True)
        )

而 Down 模块则是在池化后应用 DoubleConv:

class Down(nn.Sequential):
    def __init__(self, in_channels, out_channels):
        super(Down, self).__init__(
            nn.MaxPool2d(2, stride=2),
            DoubleConv(in_channels, out_channels)
        )

解码器与跳跃连接:细节的恢复与融合

解码器负责将瓶颈层的低分辨率特征逐步恢复至高分辨率,并通过跳跃连接融合编码器中的高分辨率细节信息。这种设计是 U-Net 能够实现精准分割的关键。

上采样与特征融合

解码器的每一层都包含一个上采样操作(如双线性插值或转置卷积),用于放大特征图尺寸。随后,将上采样结果与来自编码器同层级的特征图在通道维度上进行拼接(Concat)。拼接后,通过一个 DoubleConv 模块来融合特征并压缩通道数。

  • up1:输入 16×16×512,上采样至 32×32,与编码器 x4 拼接后,经 DoubleConv 输出 32×32×256。
  • up2:上采样至 64×64,与 x3 拼接后,输出 64×64×128。
  • up3:上采样至 128×128,与 x2 拼接后,输出 128×128×64。
  • up4:上采样至 256×256,与 x1 拼接后,输出 256×256×64。

最后,通过一个 1×1 卷积(OutConv)将 64 通道的特征图映射到目标类别数(如 num_classes=2),生成最终的像素级预测。

解码器整体结构与编码器对称,从最深层特征 x5出发,连续做 4 次上采样+拼接+双卷积,每上采样一次,H、W 乘 2;同时通道数逐层下降,最终回到与 x1 相同尺度,再用 1×1 卷积输出类别

尺寸对齐技巧

⚠️ 注意:如果输入图像的 H、W 不是 16 的整数倍,多次减半再倍增后可能会出现尺寸误差,导致无法拼接。因此,在实际实现中,通常使用 diff_x/y 计算跳跃连接分支与上采样分支的尺寸差,并通过 F.pad 进行补零操作,确保两者尺寸完全一致后再拼接。在理想的 2 的幂次方尺寸下,此操作不会产生额外开销。

转置卷积 vs 双线性插值

在 U-Net 中,上采样方式主要有两种选择:

  • 转置卷积(Transposed Convolution):其原理类似于卷积的逆操作,通过乘以转置矩阵来恢复维度。它的权重是可学习的,因此具有更强的特征恢复能力,但计算量相对较大。
  • 双线性插值(Bilinear Interpolation):一种基于距离加权的纯数学运算,计算复杂度低,但没有可学习的参数,表达能力较弱。

在代码中,通过 bilinear 参数可以灵活切换这两种上采样方式。

y = Cx\hat x = {C^T}y{H_o} = ({H_i} - 1) \times S - 2P + K + Adjf(x) \approx \frac{​{​{x_2} - x}}{​{​{x_2} - {x_1}}}f({Q_1}) + \frac{​{x - {x_1}}}{​{​{x_2} - {x_1}}}f({Q_2})Q_{11}, Q_{21}, Q_{12}, Q_{22}f({R_1}) \approx \frac{​{​{x_2} - x}}{​{​{x_2} - {x_1}}}f({Q_{11}}) + \frac{​{x - {x_1}}}{​{​{x_2} - {x_1}}}f({Q_{21}})f({R_2}) \approx \frac{​{​{x_2} - x}}{​{​{x_2} - {x_1}}}f({Q_{12}}) + \frac{​{x - {x_1}}}{​{​{x_2} - {x_1}}}f({Q_{22}})R_1, R_2f(P) \approx \frac{​{​{y_2} - y}}{​{​{y_2} - {y_1}}}f({R_1}) + \frac{​{y - {y_1}}}{​{​{y_2} - {y_1}}}f({R_2})

下面是解码器模块的代码实现,它整合了上采样、拼接和特征融合的步骤:

class Up(nn.Module):
    def __init__(self, in_channels, out_channels, bilinear=True):
        ...
        if bilinear:
            self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)
            self.conv = DoubleConv(in_channels, out_channels, in_channels // 2)
    def forward(self, x1, x2):
        x1 = self.up(x1)  # 上采样:H,W 各乘2,C不变
        diff_y = x2.size()[2] - x1.size()[2]
        diff_x = x2.size()[3] - x1.size()[3]
        x1 = F.pad(x1, [diff_x // 2, diff_x - diff_x // 2,
                        diff_y // 2, diff_y - diff_y // 2])
        x = torch.cat([x2, x1], dim=1)  # 通道拼接
        x = self.conv(x)               # DoubleConv 融合
        return x

输出层仅使用一层卷积来改变通道数,完成最终的分割图生成。

class OutConv(nn.Sequential):
    def __init__(self, in_channels, num_classes):
        super(OutConv, self).__init__(
            nn.Conv2d(in_channels, num_classes, kernel_size=1)
        )

完整网络封装与实战训练

将上述所有模块按 U-Net 的对称结构装配起来,即可得到完整的网络模型。在初始化时,需要定义输入图像通道数 in_channels、分割类别数 num_classes 以及上采样方式 bilinear。为了代码简洁,通常会定义一个基础通道数 base_c,编码器和解码器的通道数都基于它进行扩展。

    def __init__(self,
                 in_channels: int = 1,
                 num_classes: int = 2,
                 bilinear: bool = True,
                 base_c: int = 64):
        super(UNet, self).__init__()
        self.in_channels = in_channels
        self.num_classes = num_classes
        self.bilinear = bilinear

默认配置下(base_c=64),编码器通道按 64→128→256→512 递增,空间分辨率按 H→H/2→H/4→H/8→H/16 递减;解码器则完全相反,逐步恢复分辨率。这样设计出的网络参数适中,非常适合在 GPU 上进行训练。

class UNet(nn.Module):
    def __init__(self,
                 in_channels: int = 1,
                 num_classes: int = 2,
                 bilinear: bool = True,
                 base_c: int = 64):
        super(UNet, self).__init__()
        self.in_channels = in_channels
        self.num_classes = num_classes
        self.bilinear = bilinear
        self.in_conv = DoubleConv(in_channels, base_c)
        self.down1 = Down(base_c, base_c * 2)
        self.down2 = Down(base_c * 2, base_c * 4)
        self.down3 = Down(base_c * 4, base_c * 8)
        factor = 2 if bilinear else 1
        self.down4 = Down(base_c * 8, base_c * 16 // factor)
        self.up1 = Up(base_c * 16, base_c * 8 // factor, bilinear)
        self.up2 = Up(base_c * 8, base_c * 4 // factor, bilinear)
        self.up3 = Up(base_c * 4, base_c * 2 // factor, bilinear)
        self.up4 = Up(base_c * 2, base_c, bilinear)
        self.out_conv = OutConv(base_c, num_classes)

以下是 U-Net 的完整封装代码:

from typing import Dict
import torch
import torch.nn as nn
import torch.nn.functional as F
class DoubleConv(nn.Sequential):
    def __init__(self, in_channels, out_channels, mid_channels=None):
        if mid_channels is None:
            mid_channels = out_channels
        super(DoubleConv, self).__init__(
            nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=1, bias=False),
            nn.BatchNorm2d(mid_channels),
            nn.ReLU(inplace=True),
            nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1, bias=False),
            nn.BatchNorm2d(out_channels),
            nn.ReLU(inplace=True)
        )
class Down(nn.Sequential):
    def __init__(self, in_channels, out_channels):
        super(Down, self).__init__(
            nn.MaxPool2d(2, stride=2),
            DoubleConv(in_channels, out_channels)
        )
class Up(nn.Module):
    def __init__(self, in_channels, out_channels, bilinear=True):
        super(Up, self).__init__()
        if bilinear:
            self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)
            self.conv = DoubleConv(in_channels, out_channels, in_channels // 2)
        else:
            self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size=2, stride=2)
            self.conv = DoubleConv(in_channels, out_channels)
    def forward(self, x1: torch.Tensor, x2: torch.Tensor) -> torch.Tensor:
        x1 = self.up(x1)
        # [N, C, H, W]
        diff_y = x2.size()[2] - x1.size()[2]
        diff_x = x2.size()[3] - x1.size()[3]
        # padding_left, padding_right, padding_top, padding_bottom
        x1 = F.pad(x1, [diff_x // 2, diff_x - diff_x // 2,
                        diff_y // 2, diff_y - diff_y // 2])
        x = torch.cat([x2, x1], dim=1)
        x = self.conv(x)
        return x
class OutConv(nn.Sequential):
    def __init__(self, in_channels, num_classes):
        super(OutConv, self).__init__(
            nn.Conv2d(in_channels, num_classes, kernel_size=1)
        )
class UNet(nn.Module):
    def __init__(self,
                 in_channels: int = 1,
                 num_classes: int = 2,
                 bilinear: bool = True,
                 base_c: int = 64):
        super(UNet, self).__init__()
        self.in_channels = in_channels
        self.num_classes = num_classes
        self.bilinear = bilinear
        self.in_conv = DoubleConv(in_channels, base_c)
        self.down1 = Down(base_c, base_c * 2)
        self.down2 = Down(base_c * 2, base_c * 4)
        self.down3 = Down(base_c * 4, base_c * 8)
        factor = 2 if bilinear else 1
        self.down4 = Down(base_c * 8, base_c * 16 // factor)
        self.up1 = Up(base_c * 16, base_c * 8 // factor, bilinear)
        self.up2 = Up(base_c * 8, base_c * 4 // factor, bilinear)
        self.up3 = Up(base_c * 4, base_c * 2 // factor, bilinear)
        self.up4 = Up(base_c * 2, base_c, bilinear)
        self.out_conv = OutConv(base_c, num_classes)
    def forward(self, x: torch.Tensor) -> Dict[str, torch.Tensor]:
        x1 = self.in_conv(x)
        x2 = self.down1(x1)
        x3 = self.down2(x2)
        x4 = self.down3(x3)
        x5 = self.down4(x4)
        x = self.up1(x5, x4)
        x = self.up2(x, x3)
        x = self.up3(x, x2)
        x = self.up4(x, x1)
        logits = self.out_conv(x)
        return {"out": logits}

训练与评估:以视网膜血管分割为例

为了验证 U-Net 的实际效果,我们使用经典的 DRIVE 视网膜血管分割数据集进行训练。该数据集包含原始眼底图像和对应的二值分割标注,任务是区分背景与血管前景区域,因此输出类别数为 2。这是一个典型的医学图像分割任务,非常考验模型对细小结构的捕捉能力。

在训练过程中,需要配置合适的深度学习环境(如 PyTorch 和 CUDA):

python=3.10
numpy==1.22.0
pandas==1.4.4
matplotlib==3.5.3
Pillow
torch==1.13.1
torchvision==0.14.1

核心评价指标解读

为了科学评估模型性能,我们主要依赖以下指标:

  • Dice 系数:衡量预测区域与真实区域的重叠程度,取值范围 [0, 1],越高越好。
  • 像素准确率 (PA) 与平均准确率 (MPA):前者是分类正确的像素占比,后者是对每个类别计算准确率后的平均值。
  • 平均交并比 (mIoU):计算每个类别的 IoU(交集与并集之比),再取平均值,是语义分割最常用的综合性指标。
Dice(A,B) = \frac{​{2\left| {A \cap B} \right|}}{​{\left| A \right| + \left| B \right|}}Dice = \frac{​{2TP}}{​{2TP + FP + FN}}IoU = \frac{​{|A \cap B|}}{​{|A \cup B|}} = \frac{​{TP}}{​{TP + FP + FN}}

通过迭代训练,模型能够自动学习到血管的形态特征。训练完成后,我们可以观察模型在验证集上的分割效果以及各项指标的变化曲线。

结语

U-Net 凭借其对称的编码器-解码器结构和高效的跳跃连接,成功解决了语义分割中细节丢失的难题,成为该领域的基石模型。通过本文的解析,相信你已经掌握了其核心原理与 PyTorch 实现方法。无论是医学影像分析还是其他需要像素级理解的 AI 任务,U-Net 都是一个值得优先尝试的强力基线模型。

[AFFILIATE_SLOT_1]

如果你希望深入了解 Transformer 等前沿自然语言处理(NLP)技术如何与 U-Net 结合,或者想探索更多机器学习模型的优化技巧,欢迎持续关注我们的后续内容。

[AFFILIATE_SLOT_2]