ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Unet图像分割全解析:U型结构、跳跃连接与PyTorch实战

Unet图像分割全解析:U型结构、跳跃连接与PyTorch实战 “Unet网络”这四个字第一次听到的人多半会以为它和什么统一网络、联合网络沾边其实它的名字来源特别直白——整张网络左右大致对称形状像一个大写的字母U所以叫U-Net。2015年它在细胞追踪挑战赛上以明显优势拿下成绩那篇论文到今天引用量已经非常夸张。它要解决的问题也很朴素:给一张图把属于目标的那部分像素一个一个标出来也就是图像分割。医学影像里的细胞、肿瘤、视网膜血管遥感影像里的道路、建筑、水体工业质检里的划痕、瑕疵都能拿这套结构去啃。这篇我按自己当年啃Unet的顺序来写:先聊它为什么长成U形再把每层结构拆开看接着把参数计算、代码实现、训练配置一条龙走一遍最后把踩过的坑和改进方向摊开讲。不管你是刚入门语义分割的学生还是想把它落到自己业务里的工程师看完都能自己搭一版跑起来。1. Unet到底是个什么东西1.1 从一张医学图像说起假设你手里有一批显微镜下的细胞切片图每张图里密密麻麻几百个细胞挤在一起边界还互相贴着。医生想统计每个细胞的面积、数量、形态靠人工一张张描边根本不现实。这时候你要做的就是给图上每个像素打一个标签:这个像素属于细胞那个像素属于背景。这就是语义分割而Unet正是为这类任务量身的。它和分类网络最大的区别在于分类网络最后只吐出一个类别概率中间的空间信息被压没了分割网络必须保留“这个像素在哪”的信息输出一张和输入同样大小的掩膜图。你可以把分类理解成“这盘菜里有西红柿”分割则是“西红柿在这盘菜的左上角那几块区域”。Unet的结构正是围绕“既要语义、又要位置”这个矛盾设计的。语义要靠不断下采样、扩大感受野来获得位置要靠把分辨率还原回去。两条需求一拉网络自然就长成了U形。1.2 为什么传统方法不够用在深度学习普及之前分割主要靠阈值、边缘检测、区域生长、图割这些方法。它们的问题很集中:对噪声敏感、参数要手工调、换一批数据就得重新调参。细胞图像稍微有点光照不均阈值法就崩了。卷积网络出现后大家一开始的做法是“滑窗”——拿一个小窗口在图上滑动每滑一次判断中心像素的类别。这个思路准确率不差但有两个致命缺点:一是重复计算量巨大相邻窗口大量重叠区域被反复卷积二是窗口大小决定了能看到的上下文范围窗口小则上下文不足窗口大则定位精度下降。Unet用全卷积的思路一次性解决了这两个问题。整张图进去整张掩膜出来没有全连接层没有滑窗计算效率高而且通过跳跃连接把浅层的高分辨率细节和深层的语义信息结合起来定位精度和语义理解同时拿到了。1.3 U型结构的核心直觉把Unet画出来左边一列在缩小右边一列在放大中间有横线连接左右两边对应层整体像个U。左边叫编码器或者收缩路径右边叫解码器或者扩张路径中间那几条横线就是跳跃连接。编码器的活儿是“读懂图里有什么”解码器的活儿是“把这些东西画回原来的位置”跳跃连接则负责“别把原图的小细节弄丢了”。三者分工明确缺一不可。很多人第一次改Unet就喜欢动跳跃连接结果往往是细节一塌糊涂原因就在这里。2. Unet网络结构逐层拆解2.1 编码器:不断下采样把语义抓出来原版Unet的编码器一共做了四次下采样每次都是“两个3x3卷积 ReLU 一个2x2最大池化”。通道数一路翻倍:第一层64第二层128第三层256第四层512到了最底部是1024。每次下采样特征图的高宽减半通道翻倍。这个操作背后的逻辑是:随着空间分辨率降低单个神经元能“看到”的原图范围感受野在变大于是它能表达的东西从纹理、边缘逐渐升级到器官、目标这种高层概念。通道翻倍则是为了补偿空间信息损失把容量挪到通道维度上。我个人的理解是编码器在做的是一场“抽象升级”。第一层还在看像素亮度和边缘第二层看的是边角组合第三层已经能认出圆形、管状结构到了最底下那层图上某个位置激活得强不强已经和“这里是不是一团细胞”高度相关了。要注意的是原版用的是valid卷积没有padding每卷积一次特征图就缩小两圈所以最终特征图尺寸比输入小了不少。这也直接导致了输出掩膜比输入小一圈需要靠镜像补边来凑。现在复现时基本都改成padding1的same卷积尺寸对齐了代码也简单这是很常见的一个改动不算破坏原结构。2.2 解码器:一级一级把分辨率还回去解码器和编码器基本镜像。每一级先做一个上采样原版用的是2x2转置卷积把特征图高宽放大一倍、通道减半然后把编码器对应层传来的特征在通道维度上拼接再走两个3x3卷积加ReLU。这里拼接concatenate和加法add是两种不同做法。Unet原版用的是拼接通道数直接相加比如解码器某层有512通道编码器对应层也是512通道拼完就是1024通道下一层卷积再把它压回512。拼接的好处是信息保留更完整浅层特征和深层特征各占各的通道网络自己去学怎么融合加法则是把两者直接叠加参数量更省但信息融合更“粗暴”。为什么解码器要一级一级来而不是一次上采样到原尺寸因为一次放大太多细节全靠插值猜边界会糊。逐级放大、逐级融合每一级都有对应的浅层特征来“校正”边界才能清晰。2.3 跳跃连接:Unet的灵魂所在跳跃连接是Unet和普通编码器-解码器结构的分水岭。没有它网络就是一个先压缩再还原的自编码器细节丢得厉害边界一塌糊涂。它的作用是给解码器“喂”编码器里保存下来的高分辨率特征。编码器第一层的特征图分辨率最高保留了最丰富的边缘、纹理信息但语义弱解码器最后一层的特征图语义强但分辨率低。跳跃连接把两者拼起来等于让网络在做定位决策时既有“这是细胞”的高层判断又有“边界在哪个像素”的低层依据。实测里去掉跳跃连接细胞分割的边界IoU能掉十几个点肉眼看起来就是掩膜“胖了一圈”或者“缺口一堆”。所以你在改Unet时除非有非常明确的理由否则别动跳跃连接这个设计。2.4 输出层和损失函数解码器最后输出的是一个通道数等于类别数的特征图比如二分类分割就是1通道或者2通道背景前景多分类就是N通道。经过一个1x1卷积把通道压到类别数再根据任务选激活函数。二分类常用Sigmoid配合二元交叉熵多分类用Softmax配合交叉熵。医学分割里还有个经典选择是Dice损失因为前景往往只占图像的一小部分普通交叉熵会被背景主导Dice直接优化预测和真值的重叠度对小目标更友好。实际工程里我经常用“交叉熵 Dice”的加权组合两边的好处都吃训练也更稳。3. 关键参数与实现细节怎么算3.1 卷积输出尺寸的计算公式卷积输出尺寸的公式是:H_out floor((H_in 2 * padding - kernel_size) / stride) 1当kernel3、padding1、stride1时H_out H_in这就是same卷积尺寸不变。当kernel2、stride2时H_out H_in / 2这就是下采样。池化层用2x2、stride2同样让尺寸减半。转置卷积kernel2、stride2让尺寸翻倍。这些参数配合起来编码器和解码器的尺寸才能严丝合缝地对上跳跃连接才能拼接成功。3.2 特征图尺寸与通道变化全表拿输入256x256的单通道图举例走过一遍原版Unetsame卷积版本各阶段变化是这样的:阶段操作特征图尺寸通道数输入-256x2561编码1DoubleConv256x25664下采样1MaxPool 2x2128x12864编码2DoubleConv128x128128下采样2MaxPool 2x264x64128编码3DoubleConv64x64256下采样3MaxPool 2x232x32256编码4DoubleConv32x32512下采样4MaxPool 2x216x16512底部DoubleConv16x161024上采样1转置卷积拼接32x321024解码1DoubleConv32x32512上采样2转置卷积拼接64x64512解码2DoubleConv64x64256上采样3转置卷积拼接128x128256解码3DoubleConv128x128128上采样4转置卷积拼接256x256128解码4DoubleConv256x25664输出1x1卷积256x256类别数这张表我在调网络时几乎每次都要对着看尤其是自己动手改层数的时候尺寸对不上十有八九就是某一级的池化或者上采样参数写错了。3.3 参数量大致估算原版Unet同样的same卷积版本参数量大概在3100万左右其中绝大部分集中在底部那两层1024通道的卷积上。一个大致的估算是:一个3x3卷积的参数量等于 kernel_h * kernel_w * in_ch * out_ch。比如从512通道到512通道的一次3x3卷积参数就是 33512*512 ≈ 235万两个这样的卷积就是470万这还只是解码器一级里的一半。所以如果你想轻量化动刀的第一目标就是底部那些高通道层。用深度可分离卷积替换普通卷积参数量能掉到原来的八分之一到九分之一速度也快不少代价是精度可能略降一点多数任务上还能接受。3.4 评价指标怎么选分割任务最常看的是IoU也就是预测掩膜和真值掩膜的交集除以并集。它直观地反映了“预测对区域的覆盖有多准”。另一个常用的是Dice系数本质是F1在分割上的形式和IoU高度相关但数值普遍偏高一点。实际评估时我一般会同时看三类指标:整体像素准确率看大体对不对、各类别IoU看小目标有没有被忽略、边界附近的IoU看边界精不精细。只看像素准确率很容易被背景刷分比如一张图95%是背景模型全预测背景也有95%准确率但一个细胞都没分出来。4. 从零搭一个Unet:完整实操4.1 环境和数据准备环境用PyTorch就够了版本不用太纠结1.10以后的都能跑。数据这块分割任务的数据集格式灵活常见的是“原图 对应的掩膜图”成对存放掩膜是单通道的标签图像素值就是类别编号。这里有个新手特别容易踩的坑:掩膜图存成jpg再做有损压缩边界像素会被插值污染标签值变得不干净。一定要用PNG或者无损格式存掩膜。我当年就吃过这个亏训练一直在0.6的IoU上下打转换成PNG之后同样的代码直接涨到0.8以上排查了大半天才发现是格式问题。数据增强方面分割任务常用的有随机水平翻转、随机旋转、随机缩放、弹性变形医学图像尤其有用、亮度对比度扰动。注意几何变换要同时对原图和掩膜做而且掩膜必须用最近邻插值用双线性会把标签值插成小数直接出错。4.2 模型代码逐段实现先把双卷积块写出来这是Unet里重复出现次数最多的单元:import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.block nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.block(x)原版没有BatchNorm是后来大家普遍加上的。加了BN训练收敛明显更快对学习率也没那么敏感。但要注意batch很小比如2的时候BN统计量不稳这时候可以换成GroupNorm或者InstanceNorm我在小批量医学数据上基本都用InstanceNorm。下采样块和上采样块:class Down(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.pool nn.MaxPool2d(2) self.conv DoubleConv(in_ch, out_ch) def forward(self, x): return self.conv(self.pool(x)) class Up(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.up nn.ConvTranspose2d(in_ch, in_ch // 2, 2, stride2) self.conv DoubleConv(in_ch, out_ch) def forward(self, x1, x2): x1 self.up(x1) # 处理尺寸可能存在的1像素误差 diff_y x2.size(2) - x1.size(2) diff_x x2.size(3) - x1.size(3) x1 nn.functional.pad( x1, [diff_x // 2, diff_x - diff_x // 2, diff_y // 2, diff_y - diff_y // 2] ) x torch.cat([x2, x1], dim1) return self.conv(x)这里那段padding是保命的。输入尺寸不能整除16的时候编码器和解码器的特征图会出现1像素的错位直接cat会报错。加上这段自适应padding任意尺寸的输入都能跑。最后把整体拼起来:class UNet(nn.Module): def __init__(self, in_ch1, n_classes2, base64): super().__init__() self.inc DoubleConv(in_ch, base) self.d1 Down(base, base * 2) self.d2 Down(base * 2, base * 4) self.d3 Down(base * 4, base * 8) self.d4 Down(base * 8, base * 16) self.u1 Up(base * 16, base * 8) self.u2 Up(base * 8, base * 4) self.u3 Up(base * 4, base * 2) self.u4 Up(base * 2, base) self.out nn.Conv2d(base, n_classes, 1) def forward(self, x): x1 self.inc(x) x2 self.d1(x1) x3 self.d2(x2) x4 self.d3(x3) x5 self.d4(x4) x self.u1(x5, x4) x self.u2(x, x3) x self.u3(x, x2) x self.u4(x, x1) return self.out(x)把base设成64就是原版规模设成32或者16就是轻量版参数量差一个数量级适合显存紧张或者要部署到边缘设备的场景。4.3 训练流程和关键配置损失函数用交叉熵和Dice的组合这是我试下来最稳的搭配:def dice_loss(pred, target, eps1e-6): pred torch.sigmoid(pred) pred pred.view(pred.size(0), -1) target target.view(target.size(0), -1).float() inter (pred * target).sum(dim1) union pred.sum(dim1) target.sum(dim1) return 1 - ((2 * inter eps) / (union eps)).mean()优化器用Adam初始学习率1e-3或者3e-4都行配合余弦退火或者ReduceLROnPlateau。batch size在显存允许的范围内尽量大一点8到16是个舒服的区间。训练轮数看数据量几千张图的话一百来轮差不多小数据集几十张图可能几百轮就要停靠早停来卡。这里有个经验:Unet在小数据集上过拟合非常快。我做过一个只有80张图的任务训练损失一路掉到接近0验证IoU却卡在0.7不动。解决办法就是上强增强、加Dropout、加权重衰减偶尔再加点随机噪声。别指望堆数据以外的手段能彻底解决过拟合数据量上去了才是根本。4.4 推理和后处理推理时把模型切到eval模式关掉BN的统计更新和Dropout。输入归一化的方式要和训练时完全一致这点经常被忽略——训练用了均值方差归一化推理只除了255结果掉好几个点。输出经过Softmax或者Sigmoid得到概率图再取argmax或者做阈值截断得到标签图。后处理里最常用的是连通域过滤把面积过小的预测块删掉能明显减少零散噪点。如果是医学任务还可以加一点形态学开闭运算把边界磨顺。滑动窗口推理是大图必备。遥感或者病理全片动辄几千乘几千像素直接塞进网络显存扛不住得切成带重叠的块分别推理再拼回去重叠区用加权平均融合边界才不会有接缝。5. Unet的改进方向怎么选5.1 换主干网络原版编码器就是一层层堆卷积感受野和表达能力都有限。想涨点最直接的办法是把编码器换成预训练的ResNet、EfficientNet这些分类网络用它们的中间特征图接进解码器。这样能蹭到ImageNet预训练的权重小数据集上收敛快、精度高。不过换主干要注意:预训练网络的通道数和原版对不上解码器得跟着调另外预训练网络里有stride2的卷积替代了池化下采样倍率要重新核对不然拼接尺寸对不上。5.2 注意力和多尺度注意力机制在分割里用得非常广。空间注意力让网络学会“看哪里重要”通道注意力让网络学会“哪些特征通道重要”。Attention Unet就是加了注意力门控的版本在医学图像上提升明显尤其是小目标。多尺度方面有在编码器里加空洞卷积扩大感受野的有在解码器里做特征金字塔融合的还有用不同尺度的输入分别推理再融合的。思路都是让网络同时“看得近”和“看得远”。5.3 损失函数上的功夫损失函数是性价比最高的改进点改动成本低效果有时候很明显。除了交叉熵和Dice还有Tversky损失可以调FP和FN的权重、Focal损失解决类别不平衡、边界损失专门优化边界像素。类别极度不平衡的时候我一般先用Dice或Tversky压一压再配合普通的交叉熵稳一稳。5.4 轻量化让Unet跑得动如果目标是部署到移动端或者嵌入式设备轻量化是必选项。主要手段有:用深度可分离卷积替换标准卷积减少通道数用更浅的网络再配合量化。MobileNet系列的主干接进Unet是常见做法能压到几百万参数速度提升明显。代价是精度会掉掉多少取决于任务的难度。简单背景下的分割可能掉一两个点复杂场景下可能掉五个点以上。所以轻量化之前先明确精度底线别为了速度把效果做没了。6. 踩坑记录与常见问题排查6.1 常见问题速查表问题现象可能原因排查方向训练loss不降学习率过大、数据标签错乱、归一化不一致先降lr到1e-4打印几个样本的图和标签核对验证IoU远低于训练过拟合、验证集分布不同加强增强、检查数据划分是否随机预测全是背景类别极度不平衡、损失被背景主导换Dice或加权交叉熵统计前景占比边界糊成一片跳跃连接被破坏、下采样太多检查cat是否生效尝试少下采样一层显存爆了batch过大、输入尺寸过大减小batch用梯度累积或切块推理输出尺寸对不上输入尺寸非2的幂、池化上采样参数不匹配加自适应padding或resize到16的倍数6.2 小数据集上的实操心得分割标注成本很高很多项目一开始只有几十到几百张标注图这时候有几件事特别值得做。第一是把预训练编码器用上哪怕只是冻结前几层收敛也会快很多。第二是强增强弹性变形对医学图像尤其有效能显著扩大样本多样性。第三是交叉验证小数据集上单一划分的验证结果方差很大五折交叉验证出来的均值才靠谱。第四是别迷信复杂模型小数据上往往简单模型加好增强打败复杂模型。6.3 训练不收敛的排查顺序我一般的排查顺序是这样:先确认数据和标签对不对抽几张图可视化看图、掩膜、类别值是不是匹配再确认预处理训练和验证的归一化是否一致再调学习率从1e-4到1e-3之间试几个最后看模型本身用极小的子集比如4张图去过拟合如果连4张图都过拟合不了那基本就是代码或者数据有问题跟模型容量无关。这个“用几张图过拟合”的技巧非常实用能快速把问题定位到代码层还是数据层。能过拟合说明流程通了剩下的就是泛化问题过拟合不了说明流程有bug再调参也没用。6.4 使用Unet时的几个关键注意事项首先输入尺寸尽量取16的倍数因为Unet下采样四次每次减半能整除16才不会有尺寸错位。其次掩膜的插值一定用最近邻几何变换要原图和掩膜同步。第三推理阶段的预处理必须和训练完全对齐一点偏差都可能掉点。第四损失函数按任务选别一个交叉熵走天下。第五评估时不要只看像素准确率IoU和Dice才反映真实水平。第六模型checkpoint保存策略要有别只保存最后一轮验证指标最好的那轮往往才是你要的。最后分享一个我在实际项目里反复验证的体会:Unet的调参空间其实没有想象中大真正决定效果上限的是数据质量和标注一致性。同样一套代码标注干净的掩膜和随手画的掩膜训练出来的结果差得不是一星半点。所以与其反复改网络结构不如先花时间把标注规范定清楚、把数据清洗干净这部分投入的回报率远高于调网络。
返回列表