UNet++嵌套跳跃连接:缓解医学图像分割语义鸿沟
发布时间:2026/9/18 7:03:39来源:尧图网络
第一次把 UNet 的代码跑通、盯着那一堆嵌套的跳跃连接看的时候我的第一反应是这不就是在 U-Net 的跳跃连接里塞了几层卷积吗真正在医学图像分割任务上把它和 U-Net 摆在一起对比了几轮之后我才意识到这个塞卷积的动作背后藏着一个很实在的动机——编码器和解码器之间的语义鸿沟。这篇笔记不打算复述论文摘要而是把 UNet 这套嵌套 U-Net 结构拆到节点级别讲清楚它每一处设计想解决什么问题再把我复现时踩过的坑、调过的参数和验证过的结论一并摊开。如果你手上正好有小样本、边界模糊、前景占比极低的医学图像分割任务或者你在别的领域做像素级分割却一直被深层特征和浅层特征硬拼这件事困扰这篇内容应该能帮你省掉至少两轮试错。哪怕你只是想读懂这篇论文的公式和实验表格前面两节的拆解也够用了。1. 这篇 2018 年的 workshop 论文为什么被引用到现在1.1 U-Net 的瓶颈不在深度而在那条直接的跳跃连接U-Net 的结构设计在当年是相当漂亮的编码器一路下采样拿到高层语义解码器一路上采样恢复空间分辨率中间用跳跃连接把编码器的特征图直接拼到解码器对应层级上。这套结构在样本量不大的医学图像上表现非常好直到今天仍然是绝大多数分割任务的默认基线。但它的跳跃连接是一根直连的管子编码器第 i 层输出的特征图原封不动地拼到解码器第 i 层的输入里。问题就出在原封不动这四个字上。编码器浅层的特征图分辨率高、感受野小里面装的主要是边缘、角点、纹理这类低阶信息而解码器同一层级的特征图经过了一路上采样和卷积已经带上了高层语义。两者被强行拼在一起做后续卷积的时候网络得自己去调和这两种不同抽象层级的信号。在小目标、低对比度、边界糊成一片的医学图像里这种调和非常吃力。浅层特征里混着大量背景噪声和无关纹理高层语义又丢失了精细的定位信息网络往往在该听谁的话这件事上犹豫最终表现就是边界毛糙、小结构漏检、细长结构断裂。论文作者把这个现象总结为编码器与解码器特征之间的语义鸿沟semantic gap这词听起来抽象但落到实验结果上特别具体消融实验里把跳跃连接去掉指标掉得比加深网络掉得还多。1.2 用一句大白话解释语义鸿沟想象一个装修队水电工浅层特征知道每根管子的精确走向但不知道整体户型设计设计师深层特征知道哪里该做隔断、哪里该留通道但记不住每一根管子的位置。你让这两个人在同一张桌子上同时发言还不给他们翻译最后画出来的施工图必然是错位的。U-Net 的直连跳跃连接就相当于让水电工和设计师直接对话UNet 干的事情是在他们之间安排了一组逐级翻译的人——每翻译一层信息就多吸收一点对方的语境。落到网络结构上就是在跳跃路径上加卷积节点让浅层特征在被拼接之前先经过若干层卷积处理把抽象层级往上提一提从而更接近解码器那边特征的语言水平。1.3 UNet 的核心主张重新设计跳跃路径不是加深网络很多人第一次看 UNet 会误以为它是更深的 U-Net其实它的编码器主干和 U-Net 完全一样五层下采样、通道数逐层翻倍深度没有变。变的只有一处跳跃路径skip pathway从一条直连的线变成了一个带内部节点的稠密子网络。论文的原话是redesigning skip connections重新设计跳跃连接。这个定语很关键因为它决定了你后面理解一切细节的角度——UNet 的所有收益、所有代价、所有调参经验都要从跳跃路径被换了这个前提出发去推。顺带说一句作者在论文里同时给了三个配套设计嵌套的稠密跳跃路径、深度监督、以及推理阶段的剪枝。这三件事必须一起看缺一个都会让你觉得这结构又慢又没什么提升。这也是我在第三节要重点拆的内容。2. 把嵌套结构拆成节点X 的索引逻辑与计算规则2.1 两套索引i 管分辨率j 管稠密分支UNet 用了一个双下标符号 X^{i,j} 来标记网络里的每一个节点这个符号是理解整篇论文的钥匙值得花时间啃下来。i下标里的第一个数字标记节点所在的下采样层级沿编码器主干从 0 开始递增。i 相同的一排节点特征图分辨率完全相同。j第二个数字标记同一层级内沿跳跃路径的稠密分支序号从 0 开始递增。j0 就是这个层级的编码器输出节点j≥1 则是解码方向的各个节点。输出节点整个网络的最终预测放在 X^{0,4}也就是第 0 层、第 4 号分支。如果做 L3 剪枝输出就换成 X^{0,3}以此类推。以五层结构为例所有合法节点总共 15 个第 0 层有 X^{0,0} 到 X^{0,4} 共 5 个第 1 层有 X^{1,0} 到 X^{1,3} 共 4 个第 2 层 3 个第 3 层 2 个第 4 层 1 个。这个倒三角形状不是随便设计的它正好对应了每往下一层需要补的语义差就少一层这件事。2.2 节点公式的逐项拆解论文里那个公式乍一看挺唬人拆开看其实只有两条规则对于 j0 的节点纯编码器节点X^{i,0} 下采样(X^{i-1,0})就是常规的池化加卷积和 U-Net 编码器一模一样没有任何特殊之处。对于 j0 的节点X^{i,j} 卷积块( [ X^{i,0}, X^{i,1}, ..., X^{i,j-1}, 上采样(X^{i1,j-1}) ] )拆成三部分看X^{i,0} ... X^{i,j-1}同一层级里所有排在它前面的节点输出全部按通道维拼接。这就是Dense的来源——每个节点都吃到本层级的全部历史输出而不是只吃编码器那一份。上采样(X^{i1,j-1})从下一层分辨率减半那一层对应位置的节点上来先做上采样把分辨率对齐。方括号[ ]沿通道维拼接不做相加。拼接比相加保留了更多信息代价是显存和参数量上去了。注意第二项里那个j-1这是最容易写错的地方。节点 X^{i,j} 依赖的是下一层的 X^{i1,j-1}而不是 X^{i1,j}——因为下一层本来就比本层少一个节点。我第一次手写实现的时候在这里错了一格网络照样能训练、loss 照样下降但指标比论文低了三四个点排查了大半天才发现是索引的偏移错误。代码上一个通用卷积块的写法大致是这样class ConvBlock(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)以第 0 层的前三个节点为例前向过程展开是这个样子up nn.Upsample(scale_factor2, modebilinear, align_cornersFalse) # X^{0,1}吃本层编码器输出 下一层第 0 号节点上采样 x01 conv_01(torch.cat([x00, up(x10)], dim1)) # X^{1,1}同理第 1 层的第一个解码节点 x11 conv_11(torch.cat([x10, up(x20)], dim1)) # X^{0,2}吃本层所有历史节点 下一层第 1 号节点上采样 x02 conv_02(torch.cat([x00, x01, up(x11)], dim1))拼接顺序在数学上没有强制要求但一定要固定下来否则加载预训练权重的时候通道顺序会错位而且这种错误不会报错只会静默地让指标变差。2.3 手动展开 L4 的 15 个节点把依赖关系列成表会更清楚。下表按层级分组列出每个节点的输入来源节点层级输入来源说明X^{0,0}第 0 层输入图像编码器起点X^{0,1}第 0 层X^{0,0}, up(X^{1,0})最深分支的第 1 站X^{0,2}第 0 层X^{0,0}, X^{0,1}, up(X^{1,1})稠密拼接开始变长X^{0,3}第 0 层X^{0,0..2}, up(X^{1,2})L3 剪枝时的输出节点X^{0,4}第 0 层X^{0,0..3}, up(X^{1,3})默认输出节点X^{1,0}第 1 层下采样(X^{0,0})编码器第二级X^{1,1}第 1 层X^{1,0}, up(X^{2,0})解码方向X^{1,2}第 1 层X^{1,0}, X^{1,1}, up(X^{2,1})解码方向X^{1,3}第 1 层X^{1,0..2}, up(X^{2,2})只在 L4 里被用到X^{2,0}第 2 层下采样(X^{1,0})编码器第三级X^{2,1}第 2 层X^{2,0}, up(X^{3,0})解码方向X^{2,2}第 2 层X^{2,0}, X^{2,1}, up(X^{3,1})只在 L4 里被用到X^{3,0}第 3 层下采样(X^{2,0})编码器第四级X^{3,1}第 3 层X^{3,0}, up(X^{4,0})只在 L4 里被用到X^{4,0}第 4 层下采样(X^{3,0})编码器最底层最深处看这张表能发现一个有意思的事实X^{0,4} 这条完整路径实际上把所有 15 个节点都串起来了。所以你不可能靠删掉并行分支来省计算——真正能省的只有砍深度。这直接引出了第三节的剪枝话题。2.4 为什么跳跃路径的层数能补偿语义差这是整篇论文里最巧妙的一处设计值得单独说。从 X^{i,0} 出发走到 X^{i,j}中间每经过一个节点就相当于多做了一次卷积块 拼接的操作。走 4 步就是 4 个卷积块也就是 8 层 3×3 卷积。这 8 层卷积不是白加的它的作用是让原本很浅的编码器特征在逐级传递中被反复加工抽象层级一点点往上抬等到它最终被送进 X^{0,4} 的拼接操作时它和解码器那边上采样过来的特征已经处在比较接近的语义水平上了。换个说法U-Net 是让浅层特征一步到位地参与解码UNet 是让它一路走一路加工。前者靠网络自己调和后者把这个调和过程显式地建成了若干节点。这里还有个副产品。因为每个节点都同时吃到了本层的全部历史输出和下一层上来的特征所以每个节点看到的信息是多尺度的——既有本层高分辨率的位置信息也有来自更深层的语义信息。论文标题里exploit multiscale features利用多尺度特征说的就是这件事多尺度不是靠输入图像金字塔堆出来的而是从嵌套路径的结构里自然长出来的。3. 深度监督、混合损失和前向剪枝3.1 深度监督不只是多加几个输出头如果把 UNet 的四个输出节点 X^{0,1}、X^{0,2}、X^{0,3}、X^{0,4} 各自接一个 1×1 卷积加 Sigmoid再分别和金标准算损失这就是深度监督。很多人以为这只是个多加几个 head 增加损失的小技巧其实它在 UNet 里承担着两个结构性任务。第一它保证了中间分支单独拿出来也能用。剪枝之所以成立前提是 X^{0,3} 这个中间输出本身就是一个训练充分、质量可接受的预测结果。如果没有深度监督只有 X^{0,4} 收到梯度那么 X^{0,3} 只会作为中间特征存在直接拿它当输出大概率是糊的。换句话说深度监督是剪枝功能的前置条件二者是绑定的。第二它改善了梯度流动起到一定正则作用。网络里那些靠近输入的浅层节点如果只有一条很长的路径能传回梯度训练初期收敛会很慢。加了深度监督之后每个分支都有直接的监督信号梯度路径缩短训练更稳。我在小数据集上做过对照去掉深度监督头前 20 个 epoch 的验证曲线明显更抖最终 Dice 也低了大概 1.5 个点。3.2 BCE Dice 的配比与层级权重怎么定论文用的是二元交叉熵加 Dice 系数的混合损失这个组合在分割任务里几乎成了标配原因也不难理解。交叉熵是逐像素计算的每个像素一视同仁梯度稳定、收敛快但它对类别不平衡很敏感。医学图像里前景往往只占 1% 到 10%如果只用交叉熵网络只要把所有像素判成背景就能把 loss 压得很低训练直接躺平。Dice 系数是区域级别的度量它衡量的是预测区域和真实区域的重叠程度天然对前景占比低的情况更友好缺点是梯度在极端情况下不太稳分母很小的时候。两者混起来用交叉熵负责把每个像素推对方向Dice 负责把整体区域的形状拉准。我的经验配比是 1:1 起步然后看验证集上漏检多还是误检多来微调漏检多就把 Dice 的权重往上加误检尤其是小面积假阳性斑点多就把交叉熵的权重加回来一点。至于四个层级输出的权重论文里各级输出是等权求和的。我自己试过按 1、0.5、0.25 这样的比例做层级衰减越浅的分支权重越小在数据量充足的时候两者差异不明显但在样本量只有几十例的时候衰减版本明显更稳训练后期不容易出现浅层分支乱涨、把共享的低层特征带偏的情况。有一个必须提醒的细节Dice 损失里的平滑项一般取 1.0别省。医学图像里经常整张图没有前景纯背景切片如果做逐图 Dice分母会变成 0出来 NaN训练一轮就崩了。我在早期踩过这个坑加平滑项的同时还顺手过滤掉了全背景样本才算稳住。3.3 剪枝的真实含义一次训练按算力选深度这里要把前面那张依赖表再拿回来用。训练的时候全量 UNet 会计算全部 15 个节点四个输出头同时受监督。部署的时候你有两个选择继续用 X^{0,4}接受全量计算开销或者改用 X^{0,3} 当输出把 X^{0,4}、X^{1,3}、X^{2,2}、X^{3,1}、X^{4,0} 这一整块最深的结构砍掉。砍掉之后网络还剩什么第 0 到第 3 层一共 10 个节点变成一个四层深的网络。参数量、显存占用、单张推理时间全都降下来而因为深度监督已经保证了 X^{0,3} 本身训练得不错指标的下滑通常是可以接受的。这就是论文里一次训练多档速度的含义。同一次训练出来的权重你可以根据部署设备的算力在 L1 到 L4 之间挑一个档位推出不需要重新训练。对边缘设备或者需要实时处理视频流的场景这个特性比那零点几个点的 Dice 提升更有价值。论文表格里我印象比较深的几个数字是U-Net 参数量约 7.85MUNet 全量L4约 9.16M剪到 L3 之后降到 5.55M 左右推理速度和 U-Net 基本齐平。具体数值以论文原文表格为准不同实现有出入但全量比 U-Net 重、剪枝后比 U-Net 轻这个结论是稳的。3.4 别只看 Dice 那一列看这类论文的实验表格我现在的习惯是先看三列参数量、推理时间、Dice/IoU然后看第四列——验证集规模。UNet 的论文在多个数据集上做了实验包括细胞核、结肠息肉、肝脏、肺结节这几类覆盖面算是比较广的。但有几个数据集本身的测试集只有几十张图这种情况下单个 Dice 差零点几个点统计意义其实很弱。你换个随机种子重跑一遍排名的顺序都可能变。所以看表的时候心里要有个判断如果两行之间的 Dice 差距在 1 个点以内而参数量差了 30%我会更倾向于选轻的那个。这条原则在后面做模型选型的时候非常实用。4. 复现清单从 patch 采样到滑窗推理4.1 数据和增强小样本场景下的取舍医学图像分割的训练数据通常很紧张几十到几百例是常态。这时候数据策略比结构选择更重要。输入形式的选择。3D 体积CT、MRI直接上 3D UNet 显存会炸得很厉害我一般先在 2D 切片上跑通流程确认结构和损失没问题再考虑 2.5D把相邻两三张切片当成通道叠起来或者分块的 3D。2.5D 是个很好的折中既保留了 z 方向的一点上下文又不至于让显存翻三倍。patch 采样。不要整图训练尤其当原图是 512×512 甚至更大而病灶只有几十个像素的时候。按照论文里常见的做法裁成 96×96 或者 128×128 的小块。采样策略上我推荐按前景比例偏置大约 60% 到 70% 的 patch 中心落在前景上剩下的随机采。全随机采的话很多 patch 里根本看不到目标训练效率低得让人怀疑人生。增强。翻转、90 度旋转、随机缩放、弹性形变、亮度对比度抖动、高斯噪声这套组合在小样本场景下几乎是保命的。注意一点医学图像里的弹性形变幅度要控制得比自然图像更保守因为器官的解剖结构是有约束的形变太夸张等于在教网络学错误的空间关系。我一般的网格间距设在 8 到 16 像素形变强度不要超过 0.1。4.2 代码实现里最容易写错的四处索引偏移。前面提过的X^{i1,j-1}务必再核对一遍。这个错误的隐蔽之处在于它不会让程序崩溃只会让指标悄悄变差你甚至会以为是数据和超参的问题。拼接前的通道对齐。同一个层级的不同节点如果输出通道数不一致拼接操作会直接报维度错误。常见做法是同一层级内所有节点用相同的输出通道数并且在整个网络里用 32 或 64 作为基准通道数逐层翻倍。有些实现会在拼接前加一个 1×1 卷积做通道压缩这样做的好处是显存和参数量都更可控代价是多一点实现复杂度。漏掉最左边的输入节点。X^{i,j} 的稠密输入是X^{i,0}到X^{i,j-1}不是X^{i,1}到X^{i,j-1}。也就是说编码器那一份j0必须包含在内。我在早期版本里漏过这个网络依然收敛但浅层的位置信息丢失严重细长结构的断裂明显变多。上采样后的尺寸对齐。当输入尺寸不是 16 的整数倍时下采样再上采样回来的尺寸可能差一个像素拼接会报错。稳妥的做法是把输入统一 resize 或 pad 到 32 的整数倍或者在上采样后显式做一次裁剪对齐。4.3 训练配置与显存优化下面这套配置是我在多个 2D 医学分割任务上验证过、比较稳的起点你可以直接拿来当初始值项目建议值说明优化器Adam初始学习率 1e-3配合余弦退火批大小8 到 16 个 patch小 batch 时 BN 不稳考虑换 GroupNormpatch 尺寸128×128显存紧张降到 96×96基准通道32逐层翻倍到 512训练轮数150 到 300以验证集 Dice 早停为准损失BCE Dice1:1前景极稀疏时上调 Dice 权重深监督权重等权或 1/0.5/0.25小数据集建议衰减显存优化上有三招按性价比排序第一开混合精度省显存又提速几乎无副作用第二把浅层节点的通道数压住UNet 的显存大头在第 0、1 层那些高分辨率特征图上因为稠密拼接需要把它们全部保留到反向传播第三把ConvTranspose2d换成Upsample 3×3 卷积参数量少一点上采样伪影也少一点。还有一个容易忽略的点推理阶段记得把深监督头和数据流裁掉。有些实现会把四个输出头一直挂着推理时白白多算三个 1×1 卷积再加三次上采样到原图尺寸的操作对高分辨率图像来说这部分开销并不小。4.4 推理、阈值和后处理大图推理一律用滑窗。窗口尺寸和训练时的 patch 保持一致重叠比例取 0.5 左右然后用高斯权重对重叠区域的预测做加权融合——直接取平均的话窗口边缘的预测质量明显更差会在拼接处留下方格状的伪影。概率图转成二值掩码的阈值不要想当然取 0.5。在验证集上扫一遍 0.3 到 0.7把 Dice 最高的那个阈值记下来这一步往往能白捡 1 个点的提升成本几乎为零。后处理有没有用取决于你的任务。对于孤立的结节、细胞核这类目标保留最大连通域、剔除小于某个面积的小团块通常有正向收益但对于弥漫性的病灶或者细长的血管、肠壁连通域过滤会直接切碎正确结果不要用。判断标准很简单拿验证集做一次有无后处理的对照用数据说话。5. 五年后再看 UNet哪些结论站得住5.1 结构红利和训练策略红利很容易搞混论文里 UNet 相对 U-Net 的提升是实实在在的但那是在同一套训练策略下比较出来的。等你把增强、损失、学习率调度、训练轮数、测试时增强这些全都调到最优两者的差距往往会缩小。我在自己的两个数据集上做过完整对照用论文里的配置复现UNet 比 U-Net 高约 2 个点 Dice把两边的训练策略都提到同一水平更充分的增强、余弦退火、测试时翻转差距缩到 1 个点以内其中一个数据集上只有 0.4 个点。这个数量级已经接近随机种子带来的波动了。所以我对这件事的态度是UNet 那套嵌套跳跃连接确实能缓解语义鸿沟这是结构层面的真实贡献但如果你现在拿到的 U-Net 效果不好先别急着换结构把数据增强、类别不平衡处理、损失函数、学习率这几件事捋一遍收益大概率比换结构大得多。另外评估的时候一定要跑多组随机种子。只跑一次就得出某结构更好的结论在医学图像这种小数据集上非常不可靠。我现在的习惯是至少跑 3 个种子看均值和标准差差距不在一个标准差以上就不下结论。5.2 迁移到非医学场景的表现虽然论文的标题和实验都聚焦在医学图像分割这套结构的适用面其实不窄。遥感影像里的建筑提取、道路提取工业质检里的表面缺陷分割显微镜图像里的细胞计数这些任务有几个共同点目标边界不规则、正负样本极不平衡、标注数据有限。这恰好是嵌套跳跃连接擅长的场景。反过来如果你的任务是大面积规则区域的语义分割比如街景、室内场景目标往往成片出现、边界清晰U-Net 或者更轻的 DeepLab 系列已经足够UNet 多出来的那部分参数量和显存就不划算了。5.3 UNet、UNet3、nnU-Net、Transformer 系分割器怎么选这些年围绕 U-Net 的改造非常多我按自己的使用经验给一个粗略的选择逻辑U-Net数据量中等以上、任务边界相对清晰、算力受限。永远先跑这个当基线很多任务上它就能满足需求。UNet小样本、目标小而不规则、需要缓解浅层与深层特征语义差。它的剪枝特性对有推理时延要求的部署场景很友好。UNet3在 UNet 的基础上进一步做了全尺度跳跃连接每个解码节点同时吃所有尺度的编码特征和分类引导模块参数量反而更少。如果 UNet 在你任务上提升有限但显存吃紧值得试一下。自动配置类方案如 nnU-Net 那套思路真正把预处理、增强、网络配置、后处理串成流水线自动适配。它的核心观点值得记住——结构选型的影响往往小于数据处理和训练流程的影响。Transformer 系分割器在数据量充足上千例以上时能发挥出全局建模的优势但小样本场景下很容易退化而且显存和推理成本都更高。医学图像任务里我一般把它作为第二阶段尝试而不是起点。最后分享一个我在读这类论文时养成的习惯也顺手当成这篇笔记的收尾拿到任何一个新结构先别看它的 Dice 提升先看它的依赖图和参数量表画出数据从输入到输出经过的路径条数再想清楚多出来的这些路径到底在补什么信息。UNet 那 15 个节点看着复杂本质上就干了一件事——在浅层特征遇到深层特征之前多给它几次加工的机会。把这个想通了论文里的公式、深度监督、剪枝全都能顺着逻辑推出来不需要死记。
网站建设高端定制企业官网