☰
UNet图像分割全解析:从结构原理到PyTorch实战调参
2026/10/1 1:09:40 网站建设 项目流程

做图像分割这行的人,几乎绕不开 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×2561
第1阶段双卷积256×25664
池化1MaxPool 2×2128×12864
第2阶段双卷积128×128128
池化2MaxPool 2×264×64128
第3阶段双卷积64×64256
池化3MaxPool 2×232×32256
第4阶段双卷积32×32512
池化4MaxPool 2×216×16512
瓶颈层双卷积16×161024

每次池化尺寸减半、感受野翻倍,通道数翻倍则是在补偿空间信息减少带来的表达能力损失。走到瓶颈层的时候,一个 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 逐层推一遍,结果是这样的:

层操作输出尺寸感受野跳跃距离
conv1a3×3, s125631
conv1b3×3, s125651
pool12×2, s212862
conv2a/b3×3×2128142
pool22×2, s264164
conv3a/b3×3×264324
pool32×2, s232368
conv4a/b3×3×232688
pool42×2, s2167616
conv5a3×31610816
conv5b3×31614016

也就是说,瓶颈层每个特征点的感受野是 140 像素,在 256 的输入上覆盖了超过一半的边长。这就是为什么深层特征能做出可靠的语义判断,它“看到”的范围足够大。但同时也说明,一个 16×16 的瓶颈特征图,相邻两个点之间的实际间隔是 16 像素,如果你要分割的目标小于 16 像素,在瓶颈层基本就是一个点,根本无法分辨。

这个结论对实操的指导意义很直接:如果你要分割的目标普遍小于 16 像素,就不要下采样四次,改成三次,瓶颈层分辨率保持在 32×32,小目标的分割效果会好很多,代价是感受野缩小到 36 像素左右,大目标的语义判断会弱一些。这个取舍我在一个微小缺陷检测项目上验证过,目标平均尺寸只有 9 像素,改三次下采样之后 Dice 从 0.61 涨到 0.74,提升非常明显。

4.2 参数量逐层拆账

很多人只知道标准 UNet 大概 31M 参数,但不知道这些参数堆在哪里,做轻量化的时候就无从下手。我把 1 通道输入、2 类输出的版本逐层算了一遍:

模块参数量(约)占比
编码器第1阶段37 K0.1%
编码器第2阶段221 K0.7%
编码器第3阶段885 K2.9%
编码器第4阶段3.54 M11.4%
瓶颈层14.16 M45.6%
上采样层(4个)2.79 M9.0%
解码器卷积9.40 M30.3%
输出层 1×1130忽略

结论一眼就看出来了:瓶颈层一个人占了将近一半的参数。它是由 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 × 4B124 MB
梯度31M × 4B124 MB
Adam 状态31M × 4B × 2248 MB
第一层激活64 × 256 × 256 × 8 × 4B134 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,只是损失函数和增强策略调得比较细。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询