U-net 是一种对称的编码器-解码器结构的卷积神经网络 (CNN),特别适用于图像分割任务,尤其是在医学影像分割领域取得了巨大成功。U-net是2015年发的论文,在U-net网络出现之前,普遍认为深度网络的成功训练需要数千个标注训练样本。所以,U-net这篇论文提出了如何利用少样本进行深度学习,效果还很不错。由于其网络形状像“U”,故被称为U-net。
论文名字:U-net:Convolutional Networks for Biomedical Image Segmentation

图1 U-net 架构(示例为最低分辨率 32×32 像素)。每个蓝色框对应一个多通道特征图。通道数标注在框的顶部。框的左下角提供了 x-y 尺寸。白色框表示复制的特征图。箭头表示不同的操作。
一.主要贡献
1、数据增强策略:使用随机弹性变形和其他形式的数据增强来增加训练数据的多样性,从而在有限的数据集上训练出更强大的模型。
2、U 形网络结构:包含一个收缩路径(downsampling path)用于捕获上下文信息,以及一个对称的扩展路径(upsampling path)用于精确定位。
3、快速推理:U-net 能够在现代 GPU 上快速执行,对于 512×512 的图像,分割只需要不到一秒钟的时间。
4、高性能:U-net 在神经元结构分割和细胞追踪任务中表现出了卓越的性能。
二.模型结构
U-net的网络结构可以分为两个部分:左侧的编码器和右侧的解码器。编码器部分通过卷积和最大池化层提取特征,而解码器部分通过上采样和卷积层恢复图像的分辨率,并结合编码器的特征进行精细的分割。
网络结构中,编码器部分的每一层由两个卷积层和一个最大池化层组成,卷积层使用3×3的卷积核,最大池化层使用2×2的池化核。而解码器部分的每一层由一个上采样层和两个卷积层组成,上采样层使用2×2的反卷积核。如图1所示,其中:
conv 3×3,ReLu是卷积层。其中卷积核大小是3×3,再经过ReLU激活。
copy and crop是复制和裁剪。是对图片的输出尺寸进行复制并进行中心裁剪,便于与后续上采样生成的尺寸拼接。
max pool 2×2是最大池化层,卷积核为2×2。
up-conv 2×2是反卷积,用来上采样,卷积核也是2×2。
conv 1×1是卷积层,卷积核为1×1。
1.编码器部分
编码器和典型的卷积网络结构相似,它由两个3×3没有填充的卷积操作和2×2步长为2的max pooling不断重复组成。并且每个卷积操作后面都有一个ReLu激活函数。由于3×3卷积操作没有进行padding,所以每次卷积操作之后数据的宽高都会减少(k-1),k是卷积核大小。最初始的输入数据宽高为572×572,经过一次3×3没有填充的卷积之后变成了570×570。
在每次max pooling下采样中,数据的通道数会翻倍,但是宽高变为
(用于计算层与层之间的尺寸变化)。i表示输入尺寸(上一层的输出尺寸,或者最开始的输入形状),k表示卷积核大小,s表示步长。将k与s带入可以发现,每次下采样数据的高宽都会减半。
2.解码器部分
编码器中max pooling的下采样改成了步长为2的2×2的转置卷积来进行上采样。这里数据的通道数会减半,同时数据的宽高都会变为s(i-1)+k。s表示步长,i表示输入尺寸,k表示卷积核大小。将k与s带入可以发现,每次上采样数据的高宽都会翻倍。
在每次上采样之后有一个copy and crop,即跳跃连接。跳跃连接 (Skip Connections) 是一种通用的连接方式,可以将网络中某一层的输出直接传递到后面的层。 这种连接方式可以跨越任意数量的层,其主要目的是将未压缩或未经过复杂处理的特征直接传递到后面的层,以保留更多的原始信息。U-net将编码器层的特征直接传递到解码器层,能够帮助恢复图像的细节信息,即将左侧对应的特征图与上采样的输出进行concatenation。
3.其他
(1)初始化。为了确保网络各层具有近似的单位方差,U-net 使用了特殊的权重初始化策略。在具有许多卷积层和通过网络的不同路径的深度网络中,权重的良好初始化非常重要。否则,网络的某些部分可能会进行过多的激活,而其他部分永远不会起作用。理想情况下,初始化权重应该是自适应的,以使网络中的每个特征映射都具有近似的单位方差。对于具有交替卷积和ReLU层的网络,可以通过从标准偏差为
的高斯分布中绘制初始化权重来实现,N表示一个神经元传入结点的数量。这种初始化策略有助于确保网络中的每一层都能够有效地传播梯度,并且避免梯度消失或梯度爆炸的问题。
(2)数据增强。包括旋转、平移、弹性变形和灰度值变化等,其中弹性变形尤为重要,因为它能帮助网络学习到变形不变性。
(3)overlap-tile strategy (平铺策略)。原文由于当时GPU显存限制不能将原图输入,而resize会损失图像的分辨率,所以采用的是将512*512的图片进行镜像padding,得到696*696,切割出4张572*572的图片(左上,右上,左下,右下),输出388*388的图片,最后拼接在一起(重复的部分会取平均)。

图2 用于任意大图像无缝分割的重叠平铺策略(此处为 EM 堆栈中神经元结构的分割)。黄色区域的分割预测需要以蓝色区域内的图像数据作为输入。缺失的输入数据通过镜像进行外推。
(4)训练。使用较少的标注图像进行端到端的训练。采用随机梯度下降训练,由于无填充卷积,输出图像比输入少恒定的边界宽度。为了最小化开销并最大限度地利用GPU内存,在大批量数据大小的情况下使用大的输入图像块,从而将批量数据大小减少到单张图像。
三.U-net代码
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision
###卷积模块,两个conv3×3+ReLU
class conv_block(nn.Module):
def __init__(self, in_channels, out_channels, padding=0):
super().__init__()
self.conv = nn.Sequential(
nn.Conv2d(in_channels, out_channels, kernel_size=3,stride=1,padding=padding),
nn.BatchNorm2d(out_channels),
nn.ReLU(inplace=True),
nn.Conv2d(out_channels, out_channels, kernel_size=3,stride=1,padding=padding),
nn.BatchNorm2d(out_channels),
nn.ReLU(inplace=True)
)
def forward(self,x):
x = self.conv(x)
return x
###下采样,包括max pool下采样和连续的两个conv3×3+ReLU
class DownSample(nn.Module):
def __init__(self, in_channels, out_channels, padding=0):
super().__init__()
self.maxpool_conv = nn.Sequential(
nn.MaxPool2d(kernel_size=2, stride=2),
conv_block(in_channels, out_channels, padding=padding)
)
def forward(self, x):
return self.maxpool_conv(x)
###上采样,包括转置卷积上采样,并与左侧对应编码器的特征图concatenation。之后进行连续的两个conv3×3+ReLU
###解决编码器与解码器特征图尺寸不匹配的问题
class UpSample(nn.Module):
def __init__(self, in_channels, out_channels, concat=0):
super().__init__()
"""
concat=0 -> 对编码器特征图做中心裁剪
concat=1 -> 对解码器特征图做 padding 填充
concat=2 -> 在卷积块中使用 padding=1,保证卷积后尺寸不变
"""
self.concat = concat
if self.concat not in [0, 1, 2]:
raise Exception('concat not in list of [0, 1, 2]')
if self.concat == 2:
padding = 1
###反卷积上采样,将特征图宽高翻倍,通道数减半,实现上采样
self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size=2, stride=2)
###当concat=2时,padding=1,配合3×3卷积可保持特征图尺寸不变,避免裁剪或填充操作。输入通道数为in_channels(上采样结果+跳跃连接拼接后的总通道数)
self.conv = conv_block(in_channels, out_channels, padding=padding)
def forward(self, x, x_copy):
x = self.up(x)
if self.concat == 0:
B, C, H, W = x.shape
x_copy = torchvision.transforms.CenterCrop([H, W])(x_copy)
elif self.concat == 1:
diffY = x_copy.size()[2] – x.size()[2]
diffX = x_copy.size()[3] – x.size()[3]
x = F.pad(x, [
diffX // 2, diffX – diffX // 2,
diffY // 2, diffY – diffY // 2
])
###跳跃连接
x = torch.cat([x_copy, x], dim=1)
return self.conv(x)
###拼接为U-net
class UNet(nn.Module):
def __init__(self, n_channels, n_classes, concat=0):
super().__init__()
self.n_channels = n_channels
self.n_classes = n_classes
self.concat = concat
if concat == 2:
padding = 1
else:
padding = 0
expansion = 2
inplanes = 64
chns = [inplanes, inplanes * expansion, inplanes * expansion ** 2, inplanes * expansion ** 3, inplanes * expansion ** 4]
self.inc = conv_block(n_channels, chns[0], padding)
self.down1 = DownSample(chns[0], chns[1], padding)
self.down2 = DownSample(chns[1], chns[2], padding)
self.down3 = DownSample(chns[2], chns[3], padding)
self.down4 = DownSample(chns[3], chns[4], padding)
self.up1 = UpSample(chns[-1], chns[-2], concat)
self.up2 = UpSample(chns[-2], chns[-3], concat)
self.up3 = UpSample(chns[-3], chns[-4], concat)
self.up4 = UpSample(chns[-4], chns[-5], concat)
self.outc = nn.Conv2d(chns[-5], n_classes, kernel_size=1)
def forward(self, x):
e1 = self.inc(x)
e2 = self.down1(e1)
e3 = self.down2(e2)
e4 = self.down3(e3)
e5 = self.down4(e4)
x = self.up1(e5, e4)
x = self.up2(x, e3)
x = self.up3(x, e2)
x = self.up4(x, e1)
logits = self.outc(x)
return logits



