做图像分割这行的人,几乎绕不开 UNet 这个结构。我最早接触它是在一个细胞边缘分割的项目上,数据只有三十几张标注图,试了好几种网络都不太收敛,换到 UNet 之后第三天指标就起来了。后来这些年,不管是医学影像、遥感地块提取,还是工业质检里的缺陷轮廓提取,我手里跑过的分割任务里大概有七成底子都是它。UNet 图像分割这件事,说复杂也复杂,编码器解码器、跳跃连接、转置卷积、损失函数,每一样拆开都有一堆门道;说简单也简单,核心思想朴素得惊人——一边把图缩小看全局,一边把图放大定位,再把两边的信息拼回去。这篇内容我打算按自己实际搭模型、调参数的顺序来讲,从结构原理一路讲到代码实现、参数计算、显存估算、改进路线和踩过的坑,目标是让刚入门的人能照着搭出一个能跑的网络,也让已经用过 UNet 的人能从里面挑到几条之前没注意的细节。代码以 PyTorch 为例,输入尺寸、通道数、损失函数这些我都会给出具体的数字和理由,不玩虚的。
1. UNet 到底是个什么东西:从一个真实的分割任务说起
1.1 图像分割到底在解决哪三个问题
很多人第一次听到“分割”会以为是抠图,其实不完全一样。抠图是给人看的,边缘漂亮就行;分割是给下游程序看的,它要求每个像素都有一个明确的类别归属。这中间实际上有三个硬性要求同时存在,缺一个模型都会显得“不好用”。
第一是像素级分类。输入一张 512×512 的灰度图,输出也是一张 512×512 的图,只不过每个位置的数值变成了类别标签,比如 0 是背景、1 是目标。这跟分类网络输出一个向量完全不同,输出的空间维度必须保留下来。
第二是位置精度。分类任务里,目标偏个十几像素根本不影响结果;分割任务里,边界偏三个像素可能就意味着一个零件的尺寸判废。所以网络不光要知道“图里有东西”,还得知道“东西的边缘精确落在哪一行哪一列”。
第三是上下文判断。有些像素本身就长得模棱两可,比如一块阴影和一个真实的暗色区域,局部纹理几乎一样。这时候模型必须有能力看到更大范围的信息,靠周围环境来判断这个像素该归哪一类。
这三个要求互相是有冲突的。看得越全局,分辨率丢得越多,位置就越不准;盯得越局部,位置准了,但容易把噪声当目标。UNet 之所以经典,就是它用一种相当直接的方式同时照顾到了这三件事。
1.2 全卷积结构为什么在像素级任务上赢了
在 UNet 之前,很多人做分割的思路是“滑窗”——在一个大图上开一个小窗口,每次判断窗口中心那个像素属于哪一类,窗口在整个图上滑动一遍就得到完整的分割结果。这个做法逻辑上没问题,但代价大得离谱:一张 512×512 的图,如果窗口是 64×64,那就要前向推理将近 26 万次,而且相邻窗口之间大量像素是重复计算的,算力浪费严重。
全卷积网络把这件事彻底改了。它去掉网络尾部的全连接层,整个网络从头到尾都是卷积和池化,输出的就不再是一个类别向量,而是一张和输入尺寸对应的高维特征图。一次前向传播就能得到全图所有像素的预测结果,速度提升是数量级的。更关键的是,全卷积结构让网络可以接受任意尺寸的输入,训练时用什么尺寸、推理时用什么尺寸,不必严格一致,这在工程上非常友好。
我个人的体会是,全卷积这条路真正解决的是“效率”和“灵活性”两个问题,但光有全卷积还不够。早期的一些全卷积分割网络做出来,边缘都很糊,因为下采样过程中丢失的空间细节没法找回来。UNet 后面那半截上采样加跳跃连接,填的就是这个坑。
1.3 UNet 的设计动机:小样本加上精确定位
UNet 最早是为生物医学图像分割提出的,那篇论文的背景是细胞追踪挑战赛。这个场景有两个非常现实的特点。
一是标注数据极少。医学图像的标注要专业人士来做,一张图可能要花十几分钟,能拿到几十张标注图已经很不错了。这就要求网络在少量样本上也能训练,不能动辄几百万参数还依赖海量数据。
二是目标边界必须准。细胞之间常常紧挨着,两个细胞的分割结果如果粘在一起,后续的计数和形态分析就全废了。所以网络必须对边界有很强的敏感度。
UNet 的设计基本就是围绕这两点来的。它的结构是对称的 U 形,左边一路下采样把感受野做大,右边一路上采样把分辨率还原,中间用跳跃连接把同一层级的浅层特征直接送到右边。浅层特征保留了清晰的边缘和纹理,深层特征提供了语义判断,两者一拼,边界就准了。同时因为跳跃连接让梯度可以直接从解码器回传到编码器浅层,训练时的梯度流动更顺畅,小样本下也更容易收敛。这个设计在当年算是很务实的一手,没有堆什么花哨模块,就是把这几个需求串起来了。
2. UNet 网络结构逐层拆解
2.1 编码器:四次下采样里到底发生了什么
标准 UNet 的编码器由四个阶段组成,每个阶段都是“两次 3×3 卷积 + 一次 2×2 最大池化”。通道数依次是 64、128、256、512,最后到瓶颈层是 1024。假设输入是 1×256×256,走一遍下来是这样变化的:
| 阶段 | 操作 | 输出尺寸 | 输出通道 |
|---|---|---|---|
| 输入 | - | 256×256 | 1 |
| 第1阶段 | 双卷积 | 256×256 | 64 |
| 池化1 | MaxPool 2×2 | 128×128 | 64 |
| 第2阶段 | 双卷积 | 128×128 | 128 |
| 池化2 | MaxPool 2×2 | 64×64 | 128 |
| 第3阶段 | 双卷积 | 64×64 | 256 |
| 池化3 | MaxPool 2×2 | 32×32 | 256 |
| 第4阶段 | 双卷积 | 32×32 | 512 |
| 池化4 | MaxPool 2×2 | 16×16 | 512 |
| 瓶颈层 | 双卷积 | 16×16 | 1024 |
每次池化尺寸减半、感受野翻倍,通道数翻倍则是在补偿空间信息减少带来的表达能力损失。走到瓶颈层的时候,一个 16×16 的特征点实际覆盖了原图很大一块区域,语义信息非常浓缩,但空间精度已经损失了 16 倍。
这里有个容易被忽略的点:通道数翻倍不是必须的,但下采样倍数和通道数的配比会直接影响参数量分布。我见过有人把通道数改成 32、64、128、256、512 来做轻量化,参数量掉到原来的四分之一,精度在小目标上会掉得比较明显;也有人把瓶颈层加到 2048,参数量涨了一倍多,实际 Dice 只涨了不到 0.5 个点,性价比很差。
注意:池化层用最大池化还是平均池化,在分割任务里差别没有想象中大,但最大池化保留强响应、对边缘更友好,是更常见的选择。如果想进一步减少信息损失,可以用 stride=2 的卷积代替池化,让网络自己学下采样方式,代价是参数量增加。
2.2 解码器:上采样的两种做法与取舍
解码器的任务是逐步把 16×16 的特征图还原回 256×256,每一步都要把尺寸翻倍,同时把通道数减半。这里有两种主流做法,选择不同,最后的分割边缘质量和训练稳定性会有明显差别。
第一种是转置卷积,也就是 UNet 原始论文用的方式,用一个 2×2、stride=2 的卷积核来做上采样。它的好处是上采样的方式也是学出来的,理论上更灵活。但实际用起来有个很烦人的问题叫棋盘格效应:当卷积核尺寸不能被步长整除的时候,输出特征图上会出现规律的明暗格子,反映到分割结果上就是边缘出现周期性伪影。我早期在一个 PCB 缺陷分割项目上就吃过这个亏,模型训练指标看着正常,但放大看缺陷边缘有一圈锯齿,换了上采样方式之后立刻干净了。
第二种是双线性插值加普通卷积,先用插值把尺寸放大,再用一个 3×3 卷积做特征融合。这种方式没有可学习参数,不会产生棋盘格,训练也更稳。代价是表达能力弱一点,但因为后面接了卷积,实际损失不大。现在大部分工程实现我都是直接用这种方式。
两种方式的对比:
| 对比项 | 转置卷积 | 双线性插值+卷积 |
|---|---|---|
| 是否有可学习参数 | 有 | 无(后接卷积有) |
| 棋盘格伪影 | 容易出现 | 基本没有 |
| 训练稳定性 | 一般 | 好 |
| 小目标边缘表现 | 边缘略锐但可能不连续 | 边缘平滑连续 |
| 显存占用 | 略高 | 略低 |
提示:如果你必须用转置卷积,把 kernel_size 设成 stride 的整数倍(比如 2×2 配 stride 2,或 4×4 配 stride 2)能大幅减轻棋盘格问题。
2.3 跳跃连接:整个网络最值钱的那几根线
如果让我只保留 UNet 的一个设计,我会毫不犹豫选跳跃连接。它做的事情非常直接:把编码器第 1、2、3、4 阶段的输出,分别送到解码器对应阶段的输入上,在通道维度做拼接。
为什么这么关键?因为编码器在下采样过程中,浅层特征里的高频细节——边缘、纹理、细小结构——是在不断被稀释的。走到瓶颈层,语义信息很丰富,但一个细胞的两条边界可能已经被压缩到同一个像素里了,无论解码器怎么努力都还原不出来。而编码器浅层恰好保留着这些细节,直接把它们搬过来,解码器就不用从瓶颈层“凭空猜”边界在哪里。
拼接方式上有一个必须说清楚的区别:UNet 用的是通道拼接(concat),不是逐元素相加(add)。这两个操作看起来都是把两路特征合起来,但行为完全不同。相加要求两路特征通道数相同,相当于强制它们在同一语义空间里对齐,信息是“融合”的;拼接则是把两路特征并排堆在通道维度上,让后面的卷积自己去学怎么组合,信息是“并列”的。分割任务里我基本都用拼接,因为浅层和深层特征的语义差异很大,强行相加容易互相干扰。
拼接带来的直接后果是通道数翻倍。解码器第一阶段接收到的是瓶颈层上采样后的 512 通道,加上编码器第四阶段的 512 通道,一共 1024 通道。后面的双卷积就是在这个 1024 通道上做的,这也是解码器参数量的大头。
注意:拼接前一定要检查两路特征的空间尺寸是否一致。原始论文用的是无 padding 的 valid 卷积,导致编码器输出比解码器输入大,需要中心裁剪;现在大家一般都用 padding=1 保持尺寸,理论上不需要裁剪,但如果输入尺寸不是 16 的整数倍,经过四次下采样再上采样之后仍会出现 1 像素的偏差,必须做对齐处理,否则代码直接报错。
2.4 输出层与损失函数的选择逻辑
输出层很简单,就是一个 1×1 卷积,把 64 通道映射到类别数。二分类任务输出 2 个通道(背景+前景),多分类输出 N 个通道。这里要注意的是,输出的是 logits,不要在里面加 softmax,softmax 交给损失函数去做,这是 PyTorch 里的标准做法,混在一起容易导致数值不稳定。
损失函数的选择反而比输出层复杂得多。UNet 原始论文用的是带权重的交叉熵,权重是通过预先计算的距离图生成的,目的是让靠近细胞边界的像素获得更高权重,从而强迫网络关注边界。这个做法在小样本、目标紧挨的场景下效果确实好,但需要额外算权重图,实现起来麻烦。
实际工程里我用得最多的是交叉熵加 Dice 损失的组合。交叉熵逐像素算损失,收敛稳定;Dice 损失直接优化预测区域和真实区域的重叠度,对类别极度不平衡的情况(比如目标只占全图 2%)特别友好。单独用交叉熵在极端不平衡数据上很容易退化成“全预测背景”,因为全预测背景也能拿到 98% 的像素准确率,但 Dice 是 0。两个混在一起,一般给 Dice 权重 0.5 到 1.0,具体看数据不平衡程度,越不平衡 Dice 权重越高。
如果是多分类且类别之间边界特别重要,可以再加一个边界损失(Boundary Loss),它专门计算预测边界和真实边界之间的距离。我在一个视网膜血管分割任务上加过,血管的细小分支连续性明显变好,代价是训练时间增加大概两成。
3. 从零搭一个能跑的 UNet:代码级实现
3.1 环境与依赖准备
这部分没什么特别的,PyTorch 版本建议 1.10 以上,CUDA 版本跟着显卡驱动走。我平时用的组合是 Python 3.9 + PyTorch 1.13 + CUDA 11.7,比较稳。另外准备两个辅助库就够了,一个是用来读医学影像格式的,一个是用来做数据增强的。数据增强我不太推荐直接用通用图像增强库,因为分割任务里图像和掩码必须同步做几何变换,通用库处理起来容易出错,后面我会给一个自己写的增强方案。
pip install torch torchvision pip install opencv-python pip install numpy pandas matplotlib显卡方面,如果只是跑 256×256 输入、batch size 8 的二分类任务,8GB 显存就够了。如果输入上到 512×512,或者做多分类,建议 12GB 起步。这个估算我后面会详细算一遍。
3.2 编码块与解码块的具体实现
先写最基础的双卷积块。这个块会被编码器和解码器反复复用,它的设计要点是卷积后接归一化和激活。归一化这里我做了一个可切换的设计,因为医学影像数据集的 batch size 通常很小,BatchNorm 在小 batch 下方差估计不准,会出现训练和推理行为不一致的问题,这种情况我会换成 InstanceNorm。
import torch import torch.nn as nn import torch.nn.functional as F class DoubleConv(nn.Module): """两次 3x3 卷积 + 归一化 + ReLU""" def __init__(self, in_ch, out_ch, norm='batch'): super().__init__() if norm == 'batch': n1, n2 = nn.BatchNorm2d(out_ch), nn.BatchNorm2d(out_ch) else: n1 = nn.InstanceNorm2d(out_ch, affine=True) n2 = nn.InstanceNorm2d(out_ch, affine=True) self.block = nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False), n1, nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False), n2, nn.ReLU(inplace=True), ) def forward(self, x): return self.block(x)卷积层用bias=False是个小细节,因为后面紧接归一化,归一化里的 beta 参数已经承担了偏置的作用,再加偏置是冗余的,白白多几个参数。
下采样块就是池化加双卷积:
class Down(nn.Module): def __init__(self, in_ch, out_ch, norm='batch'): super().__init__() self.pool = nn.MaxPool2d(2) self.conv = DoubleConv(in_ch, out_ch, norm) def forward(self, x): return self.conv(self.pool(x))上采样块是重点。我把上采样方式和尺寸对齐逻辑都放进去了:
class Up(nn.Module): def __init__(self, in_ch, skip_ch, out_ch, up_mode='bilinear', norm='batch'): super().__init__() self.up_mode = up_mode if up_mode == 'bilinear': self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=False) conv_in = in_ch + skip_ch else: self.up = nn.ConvTranspose2d(in_ch, in_ch // 2, 2, stride=2) conv_in = in_ch // 2 + skip_ch self.conv = DoubleConv(conv_in, out_ch, norm) def forward(self, x, skip): x = self.up(x) # 尺寸对齐,处理奇数尺寸导致的 1 像素偏差 diff_y = skip.size(2) - x.size(2) diff_x = skip.size(3) - x.size(3) if diff_y != 0 or diff_x != 0: x = F.pad(x, [diff_x // 2, diff_x - diff_x // 2, diff_y // 2, diff_y - diff_y // 2]) x = torch.cat([skip, x], dim=1) return self.conv(x)那段 pad 逻辑很多人会省掉,然后输入尺寸一旦不是 16 的整数倍就报维度不匹配的错。我建议留着,花不了几个性能,省下大量调试时间。
3.3 完整模型组装与形状自检
把上面的块拼起来就得到完整网络。我加了一个base参数控制基础通道数,方便做轻量化实验:
class UNet(nn.Module): def __init__(self, in_ch=1, n_classes=2, base=64, up_mode='bilinear', norm='batch'): super().__init__() c1, c2, c3, c4, c5 = base, base*2, base*4, base*8, base*16 self.inc = DoubleConv(in_ch, c1, norm) self.down1 = Down(c1, c2, norm) self.down2 = Down(c2, c3, norm) self.down3 = Down(c3, c4, norm) self.down4 = Down(c4, c5, norm) self.up1 = Up(c5, c4, c4, up_mode, norm) self.up2 = Up(c4, c3, c3, up_mode, norm) self.up3 = Up(c3, c2, c2, up_mode, norm) self.up4 = Up(c2, c1, c1, up_mode, norm) self.outc = nn.Conv2d(c1, n_classes, 1) def forward(self, x): x1 = self.inc(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) return self.outc(x)每次写完模型我做的第一件事不是训练,而是做形状自检,用一个随机张量过一遍,确认输入输出尺寸一致、参数量符合预期:
if __name__ == '__main__': model = UNet(in_ch=1, n_classes=2, base=64) x = torch.randn(2, 1, 256, 256) y = model(x) print('输入形状:', x.shape) print('输出形状:', y.shape) total = sum(p.numel() for p in model.parameters()) print('参数量: %.2f M' % (total / 1e6))跑出来输出形状应该是(2, 2, 256, 256),参数量在 31M 左右。如果参数量明显偏小,多半是某个 Down 或 Up 块漏了;如果形状不对,检查一下跳跃连接是不是送错了层级,这种错我犯过不止一次。
3.4 训练配置的实操选择
损失函数用 Dice 加交叉熵的组合:
class DiceLoss(nn.Module): def __init__(self, smooth=1.0): super().__init__() self.smooth = smooth def forward(self, logits, target): prob = torch.softmax(logits, dim=1) n_cls = logits.shape[1] target_1h = F.one_hot(target, num_classes=n_cls) target_1h = target_1h.permute(0, 3, 1, 2).float() dims = (0, 2, 3) inter = (prob * target_1h).sum(dims) union = prob.sum(dims) + target_1h.sum(dims) dice = (2 * inter + self.smooth) / (union + self.smooth) return 1 - dice.mean() class ComboLoss(nn.Module): def __init__(self, w_dice=1.0): super().__init__() self.w = w_dice self.ce = nn.CrossEntropyLoss() self.dice = DiceLoss() def forward(self, logits, target): return self.ce(logits, target) + self.w * self.dice(logits, target)优化器我基本固定用 Adam,学习率 1e-3 起步,配合余弦退火。如果是小数据集,学习率可以降到 3e-4,避免前期震荡太厉害。batch size 在显存允许范围内尽量大,但注意如果用了 BatchNorm,batch size 小于 4 的时候要换成 InstanceNorm 或者 GroupNorm,不然统计量估计不准,训练损失降得很漂亮但验证集完全不动,这个坑我踩过。
还有个很实用的技巧是预训练权重。虽然 UNet 原版从头训也能收敛,但如果编码器换成 ResNet 之类的骨干,用 ImageNet 预训练权重初始化能明显加快收敛,在小数据集上 Dice 通常能高出两三个点。注意预训练权重大多是三通道的,如果你的输入是单通道,把第一层卷积核在通道维度求平均再复制成单通道,效果比随机初始化好很多。
4. 参数量、感受野、显存:动手前把数字算清楚
4.1 特征图尺寸与感受野的递推计算
特征图尺寸好算,每经过一次池化减半。感受野则要用递推公式:
RF_i = RF_{i-1} + (k_i - 1) * J_{i-1} J_i = J_{i-1} * s_i其中 RF 是感受野,J 是跳跃距离,初始值 RF=1、J=1,k 是卷积核尺寸,s 是步长。按标准 UNet 逐层推一遍,结果是这样的:
| 层 | 操作 | 输出尺寸 | 感受野 | 跳跃距离 |
|---|---|---|---|---|
| conv1a | 3×3, s1 | 256 | 3 | 1 |
| conv1b | 3×3, s1 | 256 | 5 | 1 |
| pool1 | 2×2, s2 | 128 | 6 | 2 |
| conv2a/b | 3×3×2 | 128 | 14 | 2 |
| pool2 | 2×2, s2 | 64 | 16 | 4 |
| conv3a/b | 3×3×2 | 64 | 32 | 4 |
| pool3 | 2×2, s2 | 32 | 36 | 8 |
| conv4a/b | 3×3×2 | 32 | 68 | 8 |
| pool4 | 2×2, s2 | 16 | 76 | 16 |
| conv5a | 3×3 | 16 | 108 | 16 |
| conv5b | 3×3 | 16 | 140 | 16 |
也就是说,瓶颈层每个特征点的感受野是 140 像素,在 256 的输入上覆盖了超过一半的边长。这就是为什么深层特征能做出可靠的语义判断,它“看到”的范围足够大。但同时也说明,一个 16×16 的瓶颈特征图,相邻两个点之间的实际间隔是 16 像素,如果你要分割的目标小于 16 像素,在瓶颈层基本就是一个点,根本无法分辨。
这个结论对实操的指导意义很直接:如果你要分割的目标普遍小于 16 像素,就不要下采样四次,改成三次,瓶颈层分辨率保持在 32×32,小目标的分割效果会好很多,代价是感受野缩小到 36 像素左右,大目标的语义判断会弱一些。这个取舍我在一个微小缺陷检测项目上验证过,目标平均尺寸只有 9 像素,改三次下采样之后 Dice 从 0.61 涨到 0.74,提升非常明显。
4.2 参数量逐层拆账
很多人只知道标准 UNet 大概 31M 参数,但不知道这些参数堆在哪里,做轻量化的时候就无从下手。我把 1 通道输入、2 类输出的版本逐层算了一遍:
| 模块 | 参数量(约) | 占比 |
|---|---|---|
| 编码器第1阶段 | 37 K | 0.1% |
| 编码器第2阶段 | 221 K | 0.7% |
| 编码器第3阶段 | 885 K | 2.9% |
| 编码器第4阶段 | 3.54 M | 11.4% |
| 瓶颈层 | 14.16 M | 45.6% |
| 上采样层(4个) | 2.79 M | 9.0% |
| 解码器卷积 | 9.40 M | 30.3% |
| 输出层 1×1 | 130 | 忽略 |
结论一眼就看出来了:瓶颈层一个人占了将近一半的参数。它是由 512→1024 和 1024→1024 两个卷积组成的,1024 通道的卷积核参数量是 512 通道的四倍,平方级增长。
所以我做轻量化的时候,第一个动的就是瓶颈层。把 1024 降到 512,参数量直接掉 7M 多,而实际精度损失在小数据集上往往不到 1 个点。反过来,如果你想加参数提升精度,往瓶颈层堆是最不划算的位置,收益最低;把钱花在第 3、4 阶段的通道数上,性价比更高。
还有一个隐藏的参数量大头是转置卷积。四个转置卷积加起来 2.79M,换成双线性插值之后这部分参数直接归零,参数量降到 28M 左右,精度基本无感。这是个很划算的替换。
4.3 显存估算与输入尺寸的经典坑
显存不够是新手最常见的问题。显存占用主要分三块:模型参数、优化器状态、中间激活值。前两块好算,第三块是大头且容易被忽略。
以输入 1×256×256、batch size 8、fp32、31M 参数的配置为例,粗略算一下:
| 项目 | 计算方式 | 占用 |
|---|---|---|
| 模型参数 | 31M × 4B | 124 MB |
| 梯度 | 31M × 4B | 124 MB |
| Adam 状态 | 31M × 4B × 2 | 248 MB |
| 第一层激活 | 64 × 256 × 256 × 8 × 4B | 134 MB |
| 全部激活(含拼接后的通道) | 约 6~8 倍第一层 | 800 MB ~ 1.1 GB |
| 合计 | - | 约 1.5 GB |
实际跑起来通常还要再多一些,因为 cuDNN 会预留工作空间,加上数据加载的显存开销,一般按理论值的 1.5 到 2 倍来预估。所以 256×256、batch 8 的配置在 4GB 卡上应该能跑,8GB 卡会比较宽裕。如果输入上到 512×512,激活值按面积算是四倍,就要 6GB 以上了。
几个省显存的实用手段:混合精度训练能省将近一半激活显存,速度还更快,我基本默认开;梯度累积可以在小显存上模拟大 batch;检查点机制能省大量激活,代价是反向传播时重算一遍,训练速度慢两到三成。这几个按需组合就行。
关于输入尺寸,有个经典坑必须说:输入的长宽最好是 16 的整数倍。因为四次下采样要除 16,如果不是整数倍,中间某一步会出现奇数尺寸,上采样之后和跳跃连接就对不齐。我上面给的 pad 逻辑能兜住,但它只是应急,长期看还是把数据统一 resize 或 pad 到 16 的倍数更省心。我习惯用 256×256 或 512×512,简单直接。
5. UNet 的改进路线:哪些真有用,哪些只是刷点
5.1 骨干网络替换:收益最稳的一类改动
把 UNet 的编码器换成更强的分类骨干,是目前最成熟的改进路线,工程上收益也最稳。常见的选择有 ResNet、EfficientNet、MobileNet、ConvNeXt 这几类。换骨干的好处有两个:一是能直接用 ImageNet 预训练权重,小数据集上收敛快;二是残差连接、深度可分离卷积这些结构本身就能提升特征质量。
具体做法是,去掉骨干网络的分类头,保留特征提取部分,然后按阶段取出四个不同分辨率的特征图接到解码器上。以 ResNet34 为例,取 layer1 到 layer4 的输出,它们的分辨率正好是输入的 1/4、1/8、1/16、1/32,通道数分别是 64、128、256、512。
这里有个细节要注意:ResNet 第一次下采样是 stride=2 的卷积,所以第一级特征就是 1/4 分辨率,比原版 UNet 少了 1/2 分辨率那一级。如果你的任务对小目标敏感,这个损失是要在意的。我一般的做法是在前面加一层轻量的 stem,把原始分辨率的信息保留下来,解码最后再拼一次。
MobileNet 那类骨干适合部署在边缘设备上,参数量能压到 5M 以内,速度提升明显,代价是小目标和细长结构的分割质量下降。我做过对比,同样数据上标准 UNet 的 Dice 是 0.847,MobileNet 版是 0.812,差了 3.5 个点,但推理速度从 42ms 降到 11ms。要不要换,取决于你的场景更看重哪个。
5.2 注意力机制与多尺度融合
注意力机制在 UNet 上的用法主要有两种。一种是在跳跃连接上加门控,让解码器决定哪些浅层特征该被采纳、哪些该被抑制,这就是 Attention U-Net 的思路。它在目标尺寸差异大的场景下有用,因为浅层特征里混着大量背景噪声,直接拼过去会干扰解码。我在一个多器官分割任务上加过,小器官的分割 Dice 涨了 4 个点,大器官几乎没变。
另一种是在编码器内部加通道注意力或空间注意力,比如 SE 模块、CBAM。这类改动的通用性更好,但提升幅度通常在 1 到 2 个点,属于锦上添花。参数量增加很少,SE 模块一般只增加百分之几,性价比还行。
多尺度融合是另一个方向,核心思路是不只用同一层级的跳跃连接,而是把更浅层的特征也引进解码器。像 UNet++ 就是把编码器的每层输出都接到多个解码器节点上,形成密集连接;UNet3+ 更进一步,把全尺度的特征都聚合起来。这类改动对小目标帮助明显,因为浅层高分辨率信息被利用得更充分,代价是显存和计算量涨得比较厉害,UNet3+ 的计算量差不多是原版的两倍。
5.3 损失函数与训练策略层面的改动
这部分改动成本最低,见效往往还快,我建议优先尝试。
Dice 和交叉熵的组合前面说过了,基本是标配。Focal Loss适合极度不平衡的场景,它会把容易分类的样本权重降下来,让模型专注难样本。不过我在分割任务上用它,效果没有在检测任务上那么惊艳,有时候反而会让边界变糊,因为难样本里包含大量噪声像素。
Tversky Loss是 Dice 的推广形式,通过两个超参数分别控制假阳性和假阴性的权重。这个在漏检代价比误检高的场景下特别好用,比如病变区域筛查,宁可多标几块也不能漏掉,把假阴性权重调高,召回率能拉上去好几个点。
边界损失专门优化边界距离,对细长结构帮助很大。我在血管、道路这类细长目标的分割上用过,连通性明显改善,断裂变少。
训练策略上,深度监督值得一提,就是在解码器的每一级输出上都加一个辅助损失。这样中间层的梯度信号更直接,收敛更快,对小数据集尤其有用。缺点是显存占用上升,因为要多算几次损失。另外课程学习也可以试,先用简单样本训,再逐步加入困难样本,训练过程更稳,但要设计样本难度的度量方式,稍微麻烦一点。
5.4 结构变体选型参考
市面上 UNet 的变体多得数不清,我把几个主流的做个对比,方便你选:
| 变体 | 核心改动 | 适用场景 | 计算量倍数 |
|---|---|---|---|
| Attention U-Net | 跳跃连接加注意力门控 | 目标尺寸差异大 | 约 1.1 |
| UNet++ | 密集跳跃连接+深监督 | 小目标、多尺度 | 约 1.5 |
| UNet3+ | 全尺度聚合+分类引导 | 复杂场景、边界要求高 | 约 2.0 |
| ResUNet | 编码器换残差块 | 通用,训练更稳 | 约 1.2 |
| TransUNet | 瓶颈层加 Transformer | 大目标、长程依赖 | 约 2.5 |
| nnU-Net | 自动配置管线 | 不想调参、要基线 | 视配置而定 |
我个人建议是:先用标准 UNet 跑一个基线,记录指标;然后按“损失函数 → 骨干网络 → 注意力 → 复杂结构”的顺序逐个试。很多时候损失函数一换指标就上来了,根本不需要动结构。反过来一上来就搞 Transformer 版本,调参调一个星期还不如基线,这种挫败感我经历过,很不值得。
nnU-Net 值得单独提一句,它不改变网络本体,而是把预处理、归一化、patch 尺寸、网络深度、损失函数这些全部自动配置一套,很多人拿它当强基线。它的配置逻辑其实挺值得学习的,比如为什么根据数据集的体素间距去决定下采样次数,为什么根据显存预算决定 patch 大小,这些思路可以直接借鉴到自己的项目里。
6. 实战踩坑与常见问题速查
6.1 数据层面的坑
类别极度不平衡是最常见的。目标只占全图百分之一二的时候,用像素准确率做评估会完全失去意义。我一般先算一下前景占比,低于 5% 就把 Dice 损失权重调高,同时在采样上做处理,让包含目标的 patch 有更高概率被抽到。做一个前景加权的采样器,比调整损失函数更直接有效。
标注质量参差不齐是另一个隐形杀手。多个人标注的掩码边界不一致,模型学出来就会很糊。我的做法是先算一下标注者之间的一致性指标,低于某个阈值的数据直接复核或者剔除。这一步看起来费时间,但它对最终指标的影响比调模型大得多。
数据增强必须图像和掩码同步。几何变换(旋转、缩放、翻转、弹性形变)必须用同一套参数作用在两者上。医学图像里弹性形变特别有用,能显著提升小数据集的泛化能力,但实现时要用一个固定的位移场同时作用,分别做两次随机形变就完全错了。颜色、亮度这类光度变换只作用在图像上,掩码不动。
注意:如果做归一化,统计量要从训练集上算,然后固定下来应用到验证集和测试集。每个样本单独做 z-score 归一化看起来方便,但会让不同样本的灰度尺度不一致,模型学得很痛苦。
6.2 训练层面的坑
验证指标不涨但训练损失一直在降,八成是过拟合。这时候先看数据量,再看模型是不是太大了。解决办法是加正则(权重衰减、dropout)、加数据增强、或者干脆把 base 通道数减半。我常用的一个粗暴办法是把瓶颈层通道从 1024 降到 512,参数量掉一半,过拟合马上缓解。
训练一开始损失就是 nan,通常是学习率太大,或者数据里有异常值。检查一下输入数据是不是有 inf 或 nan,归一化之后确认取值范围在合理区间。另外 Dice 损失里的 smooth 项要设够大,太小的话分母接近 0 会炸。
BatchNorm 在小 batch 上表现异常,前面提过了。判断方法是看训练损失和验证损失是不是差得特别远,如果是,换成 InstanceNorm 或 GroupNorm 试试。GroupNorm 的组数我一般设 8 或 16,组数太多接近 InstanceNorm,太少接近 LayerNorm。
转置卷积的棋盘格,前面也说过,换双线性插值基本解决。
双卡训练和单卡结果不一致,这种情况多半是 BatchNorm 的同步问题,多卡时要用 SyncBatchNorm。如果不确定,先把单卡跑通,指标记录下来,多卡再跑一遍对比,差得多就查归一化层。
6.3 评估与推理层面的坑
评估指标不能只看 Dice 和 IoU。这两个指标反映的是整体重叠度,对边界质量不敏感。目标的分割结果哪怕整体往内缩了一点点,边界完全错位,Dice 可能还是 0.9。所以我一般会再加一个 Hausdorff 距离(用 95 分位数版本,抗异常值),专门反映边界最大偏差。
推理时的尺寸处理和训练不一致也会导致指标暴跌。训练用 256×256 patch,推理时把整张 1024×1024 塞进去,因为分布不一致,结果会很差。正确做法是推理也用 256×256 的滑窗,重叠区域做加权融合。滑窗的步长我一般设成窗口的一半,重叠部分用高斯权重融合,边界处的拼接痕迹基本看不出来。
大图推理的拼接缝是常见问题。如果步长等于窗口大小,相邻窗口边界处的预测往往不连续,会看到明显的网格状接缝。加重叠就能解决,代价是推理时间增加。
6.4 常见问题速查表
| 现象 | 可能原因 | 处理办法 |
|---|---|---|
| 全预测背景 | 类别不平衡严重 | 提高 Dice 权重,前景加权采样 |
| 训练损失 nan | 学习率过大或数据异常 | 降 lr 到 1e-4,检查数据 |
| 边界糊、粘连 | 浅层信息利用不足 | 检查跳跃连接,试边界损失 |
| 边缘周期性锯齿 | 转置卷积棋盘格 | 改双线性插值+卷积 |
| 验证集指标不涨 | 过拟合或 lr 不对 | 加增强、减参数、调 lr |
| 小目标全部丢失 | 下采样次数过多 | 减到三次下采样 |
| 拼接前后尺寸不匹配 | 输入非 16 倍数 | resize 到 16 倍数,或加 pad |
| 多卡结果异常 | BN 同步问题 | 使用 SyncBatchNorm |
| 推理结果比训练差很多 | 输入尺寸不一致 | 保持与训练相同的滑窗推理 |
| 显存溢出 | 激活值过大 | 混合精度、梯度累积、减 batch |
最后再分享一个我在实际项目里常用的调参顺序。拿到一个新任务,先用标准 UNet 加 Dice+CE 跑一遍基线,记录 Dice、IoU、HD95 三个指标;然后把上采样换成双线性插值,看有没有变化;接着调损失函数权重,Dice 从 0.5 试到 1.5;再考虑加数据增强,尤其是弹性形变;这几步都试完还没达到要求,再动网络结构,换骨干或者加注意力。按这个顺序走,大部分任务在前三步就能达到可用水平,真正需要改结构的其实不多。我手里最近一个工业质检的项目就是这么做的,数据八百多张,最后 Dice 到 0.93,从头到尾用的都是标准 UNet,只是损失函数和增强策略调得比较细。