Unet+Resnet多尺度训练:腹部多脏器分割实战

发布时间:2026/10/11 10:43:23
Unet+Resnet多尺度训练:腹部多脏器分割实战 简介本资源面向深度学习图像分割方向的学习者与研究者提供一套基于 UnetResnet 的多尺度多类别分割实战项目及配套数据集重点解决医学图像中腹部多脏器分割的工程落地问题。压缩包共 1020 个文件以 990 张 png 图像数据为主另含 8 个 py 脚本、4 个 txt 说明、1 个 pth 权重及 readme 等整体约 363MB。项目将 Unet 的 backbone 替换为 Resnettrain 脚本会自动把数据随机缩放至设定尺寸的 0.5 至 1.5 倍实现多尺度训练utils 中的 compute_gray 函数可将 mask 灰度值保存至 txt 并自动定义输出通道预处理函数在 transforms.py 中全部重写。网络训练 50 个 epochmiou 约 0.84采用 cos 学习率衰减run_results 内附损失与 iou 曲线、训练日志及最优权重日志中可查看各类别 iou、recall、precision 与全局像素准确率预测脚本可自动推理 inference 下全部图片。目前已有 665 人学习代码注释完整参考 README 即可快速迁移到自有数据。1. 腹部多脏器分割UnetResnet 这套组合拳到底解决了什么问题腹部 CT 的多脏器分割是医学影像里最典型的「多类别、边界模糊、样本不均衡」三合一难题。肝脏、脾脏、左右肾、胰腺这五个类别在轴位切片上相邻器官的 CT 值可能只差十几个 HU胰腺还经常被胃腔和十二指肠包绕边界肉眼都难分。单靠原始 Unet 那套编码器浅层特征抓不住这种灰度差异训练到后期 Dice 卡在 0.7 上下上不去是很多人翻车的地方。把编码器换成 Resnet 预训练骨干再叠加多尺度训练本质上是让网络同时具备「强语义特征提取」和「尺度鲁棒性」两个能力。这套方案适合已经跑通过二分类分割、想往多类别医学场景推进的从业者也适合手里有腹部 CT 标注数据、想快速验证一个可复现 baseline 的团队。下面从骨干替换、多尺度策略、数据管线到避坑一步步拆开讲。2. 为什么把 Unet 编码器换成 Resnet 预训练骨干2.1 原始 Unet 编码器的三个硬伤原始 Unet 的编码器就是一层层 3x3 卷积加最大池化结构干净但在腹部 CT 这种任务上暴露三个问题。第一感受野增长慢浅层卷积堆叠再多单个像素能看到的上下文有限胰腺这种被周围器官挤压的细长结构很容易被误判成背景。第二没有预训练权重医学数据本身标注量小从零训练收敛慢且容易过拟合到训练集的灰度分布。第三梯度回传路径长深层特征在反向传播时衰减明显多类别分割里小器官比如左肾的梯度信号经常被大器官肝脏淹没。Resnet 的残差连接直接缓解了第三点恒等映射让梯度能无损回传ImageNet 预训练权重解决了第二点浅层卷积学到的边缘、纹理特征在医学图像上依然有效而 Resnet 的 stage 结构天然带来更大的有效感受野配合后续的 ASPP 或金字塔池化第一点也能补上。常见做法是把 Resnet34 或 Resnet50 的前四个 stage 拿来做编码器输出 stride 分别为 2、4、8、16 的特征图再接到 Unet 的解码器上。2.2 骨干替换的最小改动代码import torch import torch.nn as nn import torchvision.models as models class ResNetUNet(nn.Module): def __init__(self, n_classes5, backboneresnet34, pretrainedTrue): super().__init__() # 加载预训练 Resnet用 weights 参数替代旧版 pretrained weights models.ResNet34_Weights.IMAGENET1K_V1 if pretrained else None resnet models.resnet34(weightsweights) # 取前四个 stage 作为编码器 self.encoder0 nn.Sequential(resnet.conv1, resnet.bn1, resnet.relu) # stride 2 self.encoder1 nn.Sequential(resnet.maxpool, resnet.layer1) # stride 4 self.encoder2 resnet.layer2 # stride 8 self.encoder3 resnet.layer3 # stride 16 self.encoder4 resnet.layer4 # stride 32 # 解码器逐级上采样并与编码器特征拼接 self.up4 self._up_block(512, 256) self.up3 self._up_block(256, 128) self.up2 self._up_block(128, 64) self.up1 self._up_block(64, 64) self.final nn.Conv2d(64, n_classes, kernel_size1) def _up_block(self, in_ch, out_ch): return nn.Sequential( nn.ConvTranspose2d(in_ch, out_ch, kernel_size2, stride2), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, kernel_size3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x): e0 self.encoder0(x) # 1/2 e1 self.encoder1(e0) # 1/4 e2 self.encoder2(e1) # 1/8 e3 self.encoder3(e2) # 1/16 e4 self.encoder4(e3) # 1/32 d4 self.up4(e4) d3 self.up3(d4 e3) # 残差式相加也可改成 torch.cat d2 self.up2(d3 e2) d1 self.up1(d2 e1) return self.final(d1)这段代码的关键改动有三处。第一weights参数用新版 torchvision 的枚举写法旧版pretrainedTrue在新版本会报警告甚至报错这是环境配置里最常见的坑。第二编码器只取到 layer4stride 32 的特征图对 512x512 输入来说是 16x16再往下采就丢空间信息了。第三解码器里我用的是相加而不是拼接显存占用更小如果你显存充足改成torch.cat后接 1x1 卷积降维边界精度会略好一点代价是参数量增加约 30%。2.3 参数量与显存的取舍Resnet34 编码器约 21M 参数Resnet50 约 25M但 Resnet50 的 bottleneck 结构在浅层特征上不如 basic block 细腻。腹部 CT 分割里我一般先用 Resnet34 跑通Dice 稳定后再换 Resnet50 对比。输入尺寸 512x512、batch size 8 的情况下Resnet34 版本显存占用约 6.5GBResnet50 约 8.2GB单卡 12GB 以内都能跑。如果显存吃紧把输入降到 384x384或者用混合精度训练显存能再降 40% 左右。3. 多尺度训练在腹部 CT 上的具体落地方式3.1 多尺度训练到底在训练什么多尺度训练不是简单地把图片 resize 成不同尺寸喂进去它的核心目的是让网络对器官尺寸变化不敏感。腹部 CT 里肝脏在轴位上可能横跨 300 像素左肾只有 60 像素如果训练时固定一个尺度网络会偏向学习大器官的特征分布小器官的召回率明显偏低。多尺度训练通过在每个 epoch 或每个 batch 随机选择输入尺寸强迫网络在不同分辨率下都能提取到有效特征。常见做法是在 [0.75, 1.0, 1.25, 1.5] 这几个缩放比例里随机采样配合随机裁剪到统一尺寸。3.2 多尺度数据管线的代码实现import random import numpy as np import torch from torch.utils.data import Dataset import cv2 class MultiScaleAbdomenDataset(Dataset): def __init__(self, image_paths, mask_paths, base_size512, scales(0.75, 1.0, 1.25, 1.5)): self.image_paths image_paths self.mask_paths mask_paths self.base_size base_size self.scales scales def __len__(self): return len(self.image_paths) def __getitem__(self, idx): img cv2.imread(self.image_paths[idx], cv2.IMREAD_GRAYSCALE) mask cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE) # 随机选一个缩放比例 scale random.choice(self.scales) target int(self.base_size * scale) img cv2.resize(img, (target, target), interpolationcv2.INTER_LINEAR) mask cv2.resize(mask, (target, target), interpolationcv2.INTER_NEAREST) # 随机裁剪回 base_size不足则 padding if target self.base_size: x random.randint(0, target - self.base_size) y random.randint(0, target - self.base_size) img img[y:yself.base_size, x:xself.base_size] mask mask[y:yself.base_size, x:xself.base_size] else: pad self.base_size - target img np.pad(img, ((0, pad), (0, pad)), modeconstant) mask np.pad(mask, ((0, pad), (0, pad)), modeconstant) # 归一化到 [0,1] img img.astype(np.float32) / 255.0 img np.expand_dims(img, axis0) mask mask.astype(np.int64) return torch.from_numpy(img), torch.from_numpy(mask)这段代码里有两个容易忽略的细节。第一mask 的 resize 必须用INTER_NEAREST用线性插值会产生 0.5 这种非整数标签交叉熵损失直接报错。第二裁剪时如果 target 小于 base_size用 padding 补齐而不是直接 resize否则小尺度下的器官会被拉伸变形网络学到的形状先验就乱了。缩放比例我一般设四档太多档位会让 batch 内尺寸差异过大BN 层统计量不稳定。3.3 多尺度与 batch size 的配合多尺度训练时同一个 batch 里如果每张图尺寸不同没法直接堆成 tensor。两种处理方式一是整个 batch 统一用一个 scale每个 epoch 随机切换二是每张图独立随机 scale但都裁剪到同一个 base_size。我推荐第二种实现简单且尺度扰动更充分。batch size 建议不低于 8太小的话 BN 统计量波动大训练曲线会抖得厉害。如果显存只够跑 batch size 4把 BN 换成 GroupNorm稳定性会好很多。4. 五类别分割的数据准备与损失函数选择4.1 标签体系与类别不均衡腹部五类别通常是肝脏、脾脏、左肾、右肾、胰腺加上背景共六类。这五类在体积上极不均衡肝脏可能占前景像素的 60% 以上胰腺不到 5%。如果直接用交叉熵网络会倾向于把所有像素预测成肝脏和背景胰腺的 Dice 会低到 0.3 以下。常见做法是交叉熵和 Dice Loss 按 1:1 加权Dice Loss 对类别不均衡天然不敏感因为它按类别独立计算重叠度。import torch import torch.nn as nn import torch.nn.functional as F class DiceLoss(nn.Module): def __init__(self, n_classes6, smooth1e-6): super().__init__() self.n_classes n_classes self.smooth smooth def forward(self, logits, targets): probs F.softmax(logits, dim1) targets_onehot F.one_hot(targets, self.n_classes).permute(0, 3, 1, 2).float() dice 0.0 for c in range(1, self.n_classes): # 跳过背景 p probs[:, c] t targets_onehot[:, c] inter (p * t).sum() union p.sum() t.sum() dice 1 - (2 * inter self.smooth) / (union self.smooth) return dice / (self.n_classes - 1) class CombinedLoss(nn.Module): def __init__(self, n_classes6, ce_weight1.0, dice_weight1.0): super().__init__() self.ce nn.CrossEntropyLoss() self.dice DiceLoss(n_classes) self.ce_weight ce_weight self.dice_weight dice_weight def forward(self, logits, targets): return self.ce_weight * self.ce(logits, targets) self.dice_weight * self.dice(logits, targets)Dice Loss 里跳过背景是刻意的背景占比太大算进去会把前景的损失信号稀释掉。smooth设 1e-6 防止除零别设太大否则小器官的损失会被平滑掉。如果胰腺 Dice 还是上不去可以给每个类别单独设权重胰腺权重设 2.0肝脏设 0.8这个权重需要根据你数据集的统计结果调没有通用值。4.2 数据增强的边界腹部 CT 的增强不能照搬自然图像那套。水平翻转可以用因为腹部器官左右大致对称但翻转后左右肾的标签要互换这个映射关系必须写对否则网络学出来的左右肾是反的。垂直翻转慎用肝脏在上、肾脏在下的解剖关系会被破坏。旋转角度控制在 ±15 度以内大角度旋转会让器官跑到视野外。亮度对比度扰动可以加模拟不同扫描设备的灰度差异但幅度别超过 0.2否则器官间的相对灰度关系会乱。5. 训练过程中最容易翻车的几个地方5.1 预训练权重加载后 loss 不降反升现象加载 Resnet 预训练权重后前几个 epoch 的 loss 比从零训练还高。原因通常是编码器输出特征的数值范围和解码器初始化不匹配预训练权重的 BN 层统计量是基于 ImageNet 的直接冻结 BN 会导致特征分布偏移。解决办法是前 5 个 epoch 冻结编码器只训解码器让解码器先适应编码器的输出分布之后再解冻全部参数一起微调学习率降到 1e-4。5.2 多尺度训练后验证集 Dice 波动大现象训练 loss 平稳下降但验证集 Dice 每个 epoch 跳变超过 0.05。原因是多尺度训练让网络对尺度敏感而验证集如果固定一个尺度网络在非训练尺度上的表现就不稳定。解决办法是验证时也做多尺度测试把 [0.75, 1.0, 1.25] 三个尺度的预测概率图平均后再取 argmaxDice 会稳定很多代价是推理时间增加三倍。5.3 胰腺类别 Dice 始终低于 0.5现象肝脏脾脏 Dice 都在 0.9 以上胰腺卡在 0.4 到 0.5。原因是胰腺本身边界模糊且训练样本里胰腺像素太少。解决办法有三个方向一是对胰腺区域做 oversampling含胰腺的切片采样概率提高 3 倍二是在损失函数里给胰腺类别加权三是后处理阶段对胰腺预测结果做连通域分析去掉面积小于 100 像素的孤立区域。这三个手段叠加胰腺 Dice 通常能提到 0.65 以上。5.4 显存溢出但 batch size 已经降到 2现象batch size 降到 2 还是 OOM。原因往往不是 batch size而是多尺度训练时某个 scale 下的特征图太大。比如 base_size 512、scale 1.5 时实际输入是 768中间层特征图显存占用是 512 时的 2.25 倍。解决办法是给 scales 设上限最大 scale 不超过 1.25或者用梯度累积模拟大 batch累积步数设 4等效 batch size 8。5.5 验证集 Dice 高但实际预测图全是背景现象验证集 Dice 0.85但可视化预测结果发现大部分区域预测成背景。原因是背景类别占比过高Dice 计算时背景的贡献被平均进去了。解决办法是评估时只算前景类别的 Dice背景单独算 accuracy两个指标分开看。如果前景 Dice 和背景 accuracy 差距超过 0.3说明类别不均衡问题没解决好回到损失函数那一步重新调权重。6. 把多尺度推理和 TTA 叠起来榨出最后几个点训练跑通之后推理阶段还有一层提升空间。多尺度推理加测试时增强TTA是医学分割里性价比最高的后处理手段不需要重新训练只增加推理时间。具体做法是对每张测试图分别在 0.75、1.0、1.25 三个尺度下前向传播每个尺度下再做一次水平翻转总共得到 6 组 softmax 概率图。把这 6 组概率图 resize 回原始尺寸后逐像素平均再取 argmax 得到最终标签。水平翻转的概率图要记得翻回来再平均这个细节漏掉的话结果会明显变差。import torch import torch.nn.functional as F import numpy as np torch.no_grad() def multi_scale_tta_inference(model, image, scales(0.75, 1.0, 1.25)): model.eval() _, _, H, W image.shape prob_sum torch.zeros(1, 6, H, W, deviceimage.device) count 0 for scale in scales: new_h, new_w int(H * scale), int(W * scale) scaled F.interpolate(image, size(new_h, new_w), modebilinear, align_cornersFalse) # 原图预测 logits model(scaled) prob F.softmax(logits, dim1) prob F.interpolate(prob, size(H, W), modebilinear, align_cornersFalse) prob_sum prob count 1 # 水平翻转预测 flipped torch.flip(scaled, dims[3]) logits_f model(flipped) prob_f F.softmax(logits_f, dim1) prob_f torch.flip(prob_f, dims[3]) # 翻回来 prob_f F.interpolate(prob_f, size(H, W), modebilinear, align_cornersFalse) prob_sum prob_f count 1 avg_prob prob_sum / count return avg_prob.argmax(dim1)这段代码里align_cornersFalse是必须的用 True 的话插值后的概率图会有半个像素的偏移多尺度平均时边界会糊。翻转预测翻回来那一步也容易漏漏了的话翻转版本的概率图和原图对不上平均后边界反而更差。这套 TTA 在我自己的腹部五类别数据上把胰腺 Dice 从 0.62 提到了 0.68肝脏从 0.93 提到 0.94整体平均 Dice 提升约 2 个点。推理时间从单尺度单次前向的 0.3 秒涨到 1.8 秒如果做批量推理这个开销可以接受。还有一个技巧是推理时把 batch normalization 换成用训练集统计量的滑动平均而不是用当前 batch 的统计量。推理阶段 batch size 通常是 1BN 用单样本统计量会引入噪声model.eval()会自动切换到滑动平均但如果你在多尺度推理时忘了调eval()结果会差很多。这个坑我踩过训练完直接推理Dice 比验证时低了 5 个点查了半天才发现是模型还在 train 模式。最后说一个习惯每次改完骨干、损失函数或多尺度策略先在一个小验证集上跑 10 个 epoch 看趋势别一上来就训 200 个 epoch。腹部 CT 数据加载慢一次完整训练可能要大半天用小验证集快速筛掉明显不 work 的配置能省下大量时间。希望帮到你。本文还有配套的精品资源点击获取

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询