Res U-Net实战:用PyTorch复现医学图像分割模型

发布时间:2026/10/4 8:03:05
Res U-Net实战:用PyTorch复现医学图像分割模型 医学图像分割这几年已经成了深度学习落地最扎实的方向之一从病灶检测到器官勾画几乎每一家做影像AI的公司都有相关项目在跑。而U-Net这个2015年提出的结构到现在依然是医学分割领域的基线之王。不过基线归基线真要在自己的数据集上拿到更好效果很多人会优先把U-Net升级成Res U-Net——把残差连接塞进U-Net的卷积块里。这篇文章就手把手带你把Res U-Net用PyTorch完整复现一遍从网络结构拆解到训练推理直接照着写就能跑。先说清楚这篇文章适合谁刚入门医学图像分割、想搞懂U-Net变体原理的初学者已经跑通U-Net、想试试残差连接改结构的研究生以及工程开发中需要快速落地分割模型、想少踩坑的算法工程师。文章不会只贴代码会把每一步的“为什么这么做”拆开讲透毕竟理解了原理改起结构来才顺手。1. Res U-Net网络结构拆解到底改了什么1.1 从U-Net到Res U-Net的演进逻辑U-Net之所以在医学图像分割里经久不衰核心在于它的编码器-解码器结构和跳跃连接skip connection。编码器逐层下采样提取语义特征解码器逐层上采样恢复空间分辨率跳跃连接把同尺度的低层细节特征直接拼到高层语义特征上弥补下采样丢失的边界信息。这个设计特别契合医学图像的痛点病灶区域小、边界模糊、标注样本少低层特征和高层特征的融合能让分割结果更精细。但U-Net也有一个绕不开的问题网络一深梯度消失和退化现象就来了。残差网络ResNet当年在ImageNet上大杀四方靠的就是残差连接解决了深度网络的训练难题。Res U-Net的思路很直接——把U-Net里的每一个卷积块从“卷积-激活-卷积”换成带残差分支的残差块。这样网络加深了反向传播时梯度能沿着残差捷径直接流回浅层训练稳定性大幅提升。这里有个容易混淆的点U-Net本身已经有跳跃连接了Res U-Net的残差连接和它不是一回事。跳跃连接连接的是编码器和解码器是跨层级的残差连接连接的是同一卷积块内的输入和输出是“块内”的。两者并不冲突反而是互补的——块内残差保证梯度畅通跨层级联跳保证特征融合。理解了这层关系你就明白Res U-Net为什么能同时兼顾深度和细节了。1.2 ResBlock核心实现不要照抄ResNet的瓶颈结构实现Res U-Net时最常见的误区是直接把ResNet里的Bottleneck瓶颈块搬过来用。但医学图像分割任务输入的原始图像分辨率普遍偏高512x512甚至更高通道数相对较少直接套用1x1卷积降维的瓶颈结构反而会丢失细节信息。更合理的做法是沿用U-Net中的DoubleConv思路将两个3x3卷积组合成一个残差块。我把核心的ResBlock代码贴出来这是整个Res U-Net的地基import torch import torch.nn as nn class ResBlock(nn.Module): def __init__(self, in_channels, out_channels, stride1): super(ResBlock, self).__init__() self.conv1 nn.Conv2d(in_channels, out_channels, kernel_size3, stridestride, padding1, biasFalse) self.bn1 nn.BatchNorm2d(out_channels) self.relu nn.ReLU(inplaceTrue) self.conv2 nn.Conv2d(out_channels, out_channels, kernel_size3, stride1, padding1, biasFalse) self.bn2 nn.BatchNorm2d(out_channels) # 残差分支当输入输出通道不一致时用1x1卷积对齐 self.shortcut nn.Sequential() if stride ! 1 or in_channels ! out_channels: self.shortcut nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size1, stridestride, biasFalse), nn.BatchNorm2d(out_channels) ) def forward(self, x): identity self.shortcut(x) out self.conv1(x) out self.bn1(out) out self.relu(out) out self.conv2(out) out self.bn2(out) out identity out self.relu(out) return out几个关键细节值得单独拎出来说。第一第一个卷积层设置了stride参数这是为了在下采样时同时完成空间分辨率的降低和通道数的扩张。第二shortcut分支只在输入输出通道不一致或stride不为1时才生效避免无谓的计算开销。第三经过shortcut的1x1卷积后BatchNorm也是必须的否则两个分支的数值分布不对齐相加后特征分布会被破坏。提示关于biasFalse这个设置可能是因为2D卷积后跟了BatchNormBatchNorm自身带有偏置项卷积层再设置biasTrue会造成冗余的参数量。2. 编码器与解码器设计整体架构怎么搭2.1 Encoder特征提取逐层下采样的节奏控制Res U-Net的编码器和标准U-Net一样共有4次下采样每次下采样后特征图的通道数翻倍。按照经典的U-Net配置初始通道数设为64经4次下采样后通道数依次变为64、128、256、512。每个下采样阶段由一个stride为2的ResBlock完成在减少空间尺寸的同时扩张通道数。class Encoder(nn.Module): def __init__(self, in_channels1, init_features64): super(Encoder, self).__init__() # 第一层不做下采样先做特征提取 self.conv1 ResBlock(in_channels, init_features, stride1) # 后续四层逐层下采样并增加通道数 self.conv2 ResBlock(init_features, init_features * 2, stride2) self.conv3 ResBlock(init_features * 2, init_features * 4, stride2) self.conv4 ResBlock(init_features * 4, init_features * 8, stride2) self.conv5 ResBlock(init_features * 8, init_features * 16, stride2) def forward(self, x): # 保存每一层的输出用于跳跃连接 features [] x self.conv1(x) features.append(x) x self.conv2(x) features.append(x) x self.conv3(x) features.append(x) x self.conv4(x) features.append(x) x self.conv5(x) features.append(x) return features编码器的输出是一个特征列表把每一层的特征图都保存下来。注意最后一层conv5的stride也是2也就是说整个网络实际上做了4次真正的下采样特征图从原始分辨率一路缩到1/16。这个下采样倍数兼顾了感受野和分辨率之间的平衡既与大病灶的全局语义有关也不至于让细节丢失太多。2.2 Decoder上采样恢复跳跃连接为什么不能直接拼接解码器的作用是把编码器输出的低分辨率特征逐步恢复到原始分辨率。每一层解码器做的事是先对上一层的特征图进行2倍上采样然后把上采样结果与编码器对应层的特征图在通道维度上拼接最后经过一个ResBlock进行特征融合。class Decoder(nn.Module): def __init__(self, features): super(Decoder, self).__init__() self.up1 nn.ConvTranspose2d(features[4], features[3], kernel_size2, stride2) self.res1 ResBlock(features[4], features[3]) self.up2 nn.ConvTranspose2d(features[3], features[2], kernel_size2, stride2) self.res2 ResBlock(features[3], features[2]) self.up3 nn.ConvTranspose2d(features[2], features[1], kernel_size2, stride2) self.res3 ResBlock(features[2], features[1]) self.up4 nn.ConvTranspose2d(features[1], features[0], kernel_size2, stride2) self.res4 ResBlock(features[1], features[0]) def forward(self, features): x features[-1] # 逐层上采样并与编码器特征拼接 x self.up1(x) x torch.cat([x, features[3]], dim1) x self.res1(x) x self.up2(x) x torch.cat([x, features[2]], dim1) x self.res2(x) x self.up3(x) x torch.cat([x, features[1]], dim1) x self.res3(x) x self.up4(x) x torch.cat([x, features[0]], dim1) x self.res4(x) return x很多人不理解为什么上采样后还要做一次拼接再过一个ResBlock为什么不直接卷积输出。原因在在于上采样无论是转置卷积还是插值生成的是一种“稀疏”的特征分布直接和编码器特征拼接会导致数值尺度不一致。再过一个ResBlock可以让解码器充分学习到“如何把低层细节和高层语义融合在一起”而不是死板地把两类特征简单叠放。这种设计的灵活性正是Res U-Net在多个医学分割数据集上表现稳定的原因之一。3. 完整Res U-Net组装与关键参数选择3.1 主网络组装从输入到分割输出的完整流程把编码器和解码器拼起来还需要在最后接一个输出头。医学图像分割本质上是一个逐像素分类问题对于二分割任务如分割病灶区域输出层用1x1卷积把通道数降为1再接Sigmoid得到每个像素属于前景的概率图。对于多类别分割任务则把通道数设为类别数接Softmax。class ResUNet(nn.Module): def __init__(self, in_channels1, out_channels1, init_features64): super(ResUNet, self).__init__() self.encoder Encoder(in_channels, init_features) features [init_features, init_features * 2, init_features * 4, init_features * 8, init_features * 16] self.decoder Decoder(features) self.final nn.Conv2d(init_features, out_channels, kernel_size1) def forward(self, x): features self.encoder(x) x self.decoder(features) x self.final(x) return x整个网络的可训练参数大约在3300万到3900万之间具体取决于初始通道数的设置。以init_features64为例模型大小约130MB左右float32精度在常见的12GB显存GPU上batch size设为8、输入尺寸512x512可以正常训练。如果显存吃紧可以把输入尺寸降到256x256或把初始通道数降到32参数量会减少四倍多对精度的影响相对有限。3.2 初始通道数对性能的影响怎么选合适初始通道数是Res U-Net最敏感的超参数之一。通道数翻倍直接带来参数量和计算量的4倍增长但并非越大越好。从我的实测经验看初始通道数参数量约显存占用512x512输入适用场景16210万约2GB小数据集、单张医疗影像快速验证32830万约5GB中等数据量、追求训练效率643300万约11GB标准配置效果和资源消耗最均衡1281.3亿显存需求过高很少用除非有大显存且数据量足够对于大多数医学分割任务64是甜点值。如果数据集只有几百张甚至更少建议降到32不然很容易过拟合。判断标准其实很简单用训练集和验证集的Dice系数曲线做对比如果训练Dice不断上升但验证Dice停滞甚至下降基本就是模型容量过大、数据量不够学。4. 训练策略与损失函数实战让网络真正收敛4.1 医学分割的损失函数选型BCE还是Dice Loss复现网络只是第一步真正决定分割效果的是训练策略。医学图像分割中正负样本比例严重失衡是常态——病灶区域往往只占整张图像的1%甚至更少。如果直接用普通的交叉熵损失BCE Loss网络会倾向于把所有像素都预测为背景因为这样做准确率已经很高了。Dice Loss是解决类别不平衡问题的利器。它直接优化Dice系数对前景和背景的数量差异不敏感是医学分割任务中的标配选择。训练时也可以把BCE Loss和Dice Loss结合组成混合损失函数class BCEDiceLoss(nn.Module): def __init__(self, weight_dice0.5): super(BCEDiceLoss, self).__init__() self.bce nn.BCEWithLogitsLoss() self.weight_dice weight_dice def forward(self, logits, targets): bce_loss self.bce(logits, targets) probs torch.sigmoid(logits) smooth 1e-5 # 将概率图展平便于计算整个batch的Dice probs probs.contiguous().view(probs.size(0), -1) targets targets.contiguous().view(targets.size(0), -1) intersection (probs * targets).sum(dim1) dice (2.0 * intersection smooth) / (probs.sum(dim1) targets.sum(dim1) smooth) dice_loss 1.0 - dice.mean() return bce_loss self.weight_dice * dice_loss这段代码里有几个细节需要留意。BCEWithLogitsLoss接收的是未经过Sigmoid的logits这样在数值上更稳定避免了Sigmoid后再计算交叉熵导致的梯度消失问题。Dice Loss计算时给分子分母都加了平滑项smooth防止出现0/0的情况。所以这里的weight_dice参数一般取0.5到1.0之间太大会导致训练早期梯度震荡太剧烈。4.2 训练过程中的数据增强与预处理医学图像数据增强和自然图像有些区别。最基本的原则是不能随意改变图像的解剖结构对应关系。像随机裁剪、翻转、旋转、缩放这些几何变换是安全的而色彩抖动、随机擦除这些在分类任务中常用的增强手段在医学分割中要非常谨慎使用因为灰度值的细微变化可能影响病灶的可见性。我常用的医学分割增强组合是随机水平翻转、随机旋转范围±15度、随机缩放0.9到1.1倍、弹性形变。弹性形变对医学图像特别有效因为人体组织本身就有一定形变自由度能有效提升模型的泛化能力。PyTorch里可以用albumentations库快速实现import albumentations as A from albumentations.pytorch import ToTensorV2 train_transform A.Compose([ A.RandomRotate90(p0.5), A.HorizontalFlip(p0.5), A.VerticalFlip(p0.5), A.ElasticTransform(alpha1, sigma50, p0.3), A.Resize(256, 256), A.Normalize(mean(0.5,), std(0.5,)), ToTensorV2(), ]) test_transform A.Compose([ A.Resize(256, 256), A.Normalize(mean(0.5,), std(0.5,)), ToTensorV2(), ])关于Normalize医学图像如CT、MRI的像素值范围差异很大CT是Hounsfield单位通常范围在-1000到3000多MRI则没有统一的绝对尺度。实践中一般先做MinMax归一化或Z-score标准化再在训练时做统一Normalize。mean0.5、std0.5是假设数据已经归一化到0到1区间后的常见设置实际使用时应根据自己的数据分布调整。4.3 优化器与学习率调度让网络稳定收敛Res U-Net的训练推荐使用Adam优化器初始学习率设为1e-4到3e-4之间。很多初学者会直接套用自然图像分类任务中的初始学习率如1e-3这在医学分割任务中往往导致训练早期就出现梯度爆炸。原因在于分割任务的损失函数尤其是Dice Loss在训练初期梯度变化非常剧烈学习率过大会让权重更新幅度失控。学习率调度器我推荐使用ReduceLROnPlateau当验证集Dice系数连续若干轮不再提升时把学习率降低为原来的1/2或1/5。这个策略比固定步长衰减更稳定因为不同数据集上模型收敛的速度差异很大固定步长很难一次调好。optimizer torch.optim.Adam(model.parameters(), lr1e-4, weight_decay1e-5) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemax, factor0.5, patience10, verboseTrue ) # 每个epoch结束时 scheduler.step(val_dice)关于weight_decay我习惯设一个很小的值1e-5主要是约束大权重防止模型对训练集产生过拟合。但不要设太大医学分割数据集规模有限正则化太强会损失对细节边界的分割能力。5. 数据集准备与评估指标复现效果怎么量化5.1 数据集格式与DataLoader封装医学图像分割最经典的开源数据集之一就是2015年ISBI的细胞膜分割数据集DSB2015以及后续的脑肿瘤数据集BraTS。这里以二分割任务为例演示如何把扫描图像和对应的掩膜Mask封装成PyTorch的Dataset类。一个容易踩的坑是医学图像的原始格式通常是DICOM、NIfTI.nii.gz等并非直接读入就能用。DICOM需要处理窗宽窗位NIfTI则要处理方向矩阵。这里建议先用SimpleITK或NiBabel把数据转为PNG或NPY格式做预处理训练时直接读处理好的文件能省掉很多运行时的问题。下面是一个简单的标准数据加载流程import os import numpy as np import torch from torch.utils.data import Dataset from PIL import Image class MedicalSegDataset(Dataset): def __init__(self, image_dir, mask_dir, transformNone): self.image_dir image_dir self.mask_dir mask_dir self.transform transform self.images sorted(os.listdir(image_dir)) def __len__(self): return len(self.images) def __getitem__(self, idx): img_path os.path.join(self.image_dir, self.images[idx]) mask_path os.path.join(self.mask_dir, self.images[idx].replace(.npy, _mask.npy)) image np.load(img_path).astype(np.float32) mask np.load(mask_path).astype(np.float32) # 统一扩展成单通道格式 image image[None, ...] # (1, H, W) mask mask[None, ...] # (1, H, W) return torch.from_numpy(image), torch.from_numpy(mask)这里直接用了NPY格式进行演示实际项目中如果数据是PNG格式用PIL读取后再转Tensor即可。Dataset封装的关键点是所有数据在进入网络之前尺寸、通道数、数据类型都要保持一致否则如果在训练中途才报错排查起来非常浪费时间。5.2 分割指标计算Dice、IoU、HD95怎么算医学图像分割的评估指标和自然图像分割差别较大。语义分割常用的mIoU忽略了病灶区域的体积占比在医学任务中参考价值有限。医学领域最看重的是Dice系数和95%豪斯多夫距离HD95。Dice关注重叠程度HD95关注边界偏差两者结合才能全面评估分割质量。Dice系数的公式是2乘以预测和真实标签的交集面积除以两者面积之和。在代码实现时我一般会分别计算每个样本的Dice再取平均而不是在整个batch上直接计算因为后者会被大病灶样本主导掩盖小病灶样本上的表现。def compute_dice(preds, masks, threshold0.5): preds torch.sigmoid(preds) preds (preds threshold).float() intersection (preds * masks).sum(dim(1, 2, 3)) union preds.sum(dim(1, 2, 3)) masks.sum(dim(1, 2, 3)) dice (2.0 * intersection 1e-6) / (union 1e-6) return dice.mean().item()HD95的计算相对复杂需要用到scipy.ndimage的距离变换函数。它的含义是预测边界和真实边界之间有95%的点距离在某个阈值以内这个值越小说明预测边界的精度越高。在实现时可以用scipy.ndimage.distance_transform_edt计算真实掩膜到背景的距离图再统计预测边界上这些距离值的95%分位数。我在训练时一般每5个epoch就在验证集上完整跑一次Dice和HD95记录到日志里。相比只看训练loss下降这两个指标更直观地反映分割质量的变化趋势。6. 训练与推理常见问题排查实录6.1 训练不收敛或loss震荡从这些方向排查Res U-Net训练中遇到的最典型问题就是loss震荡不下降。我在实际项目中踩过的坑基本集中在以下三个方面第一数据没有归一化好。医学图像的对比度差异非常大如果你直接把读入的原始像素值丢进网络卷积计算出的特征图数值范围会非常不稳。检查方法很简单训练前打印一下输入数据的min、max、mean、std如果范围还是上千上万说明预处理步骤有问题。第二Dice Loss的权重设置不合理。如果dice_loss权重太高比如大于1训练初期梯度震荡会比较剧烈表现为loss曲线忽上忽下。建议先只使用BCE Loss训练10个epoch让网络稳定下来再切换到BCEDice混合损失微调这种两阶段训练法在医学分割里非常实用。第三BatchNorm对batch size敏感。如果显存有限导致batch size只能设为2或者更小BatchNorm统计量会非常不稳定。这种情况下建议使用GroupNorm替代BatchNorm或者使用梯度累积技巧来模拟更大的有效batch size。6.2 Windows环境常见的OSError 1114报错在Windows上用PyTorch复现医学图像分割模型时一个非常常见的报错是OSError: [WinError 1114] 动态链接库(DLL)初始化例程失败。 Error loading C:\Users\xxx\.conda\envs\pytorch\lib\site-packages\torch\lib\c10.dll or one of its dependencies.这个问题的本质是系统缺少必要的C运行库或者PyTorch版本和当前的CUDA驱动不匹配。我的排查步骤是先安装最新的Visual C Redistributable运行库这是最常见的解决方案然后检查显卡驱动版本是否支持已安装的CUDA版本如果驱动太旧直接安装CPU版本的PyTorch测试最后一招是用conda install而不是pip install安装PyTorchconda会顺带安装OpenMP等依赖库通常能解决大部分DLL报错问题。6.3 推理阶段小目标分割效果差别急着加网络深度训练完成的模型总会在某些样本上表现不佳尤其是小病灶或者边界模糊的区域。我建议按以下顺序排查先检查数据预处理是否引入了偏差比如归一化方式对灰度范围较窄的医学图像是否合适再看数据增强是否过强导致病灶形态被改变失真最后才考虑网络结构问题。实际上对于小目标分割增加网络中低层特征的通道数或者增强跳跃连接的作用往往比无脑加深网络更有效。从实际经验来看如果模型在验证集上的Dice已经达到85%以上还想继续提升优先尝试的应该是SWA随机权重平均、TTA测试时增强、EMA指数移动平均这些训练技巧而不是动网络结构。TTA对医学分割的提升尤其明显推理时将图像做水平翻转和垂直翻转分别预测后取概率平均再用阈值得到最终分割结果通常能提升0.5到1个百分点的Dice。7. 复现之后还能怎么玩Res U-Net的变体与扩展思路把Res U-Net跑通只是迈出了一小步。在实际项目里你会发现Res U-Net依然有不少可以优化的空间。与Attention U-Net结合在跳跃连接处加入注意力门控机制可以抑制无关背景区域的特征响应与Dense U-Net结合用密集连接替代残差连接能进一步提升特征复用效率代价是参数量和显存占用同步上升。在三维医学图像如CT、MRI的体数据上可以把Res U-Net中的2D卷积和池化全部换成3D版本处理带有空间上下文信息的数据。3D模型对显存的要求成倍增加实际落地时通常用2.5D策略——在三个正交平面上分别用2D Res U-Net做分割然后把三个方向的预测结果融合。这种方案兼顾了3D上下文信息和2D模型的高效性是目前很多医疗AI公司的工程选择。从我在多个医学分割竞赛和实际临床项目中的体会来看Res U-Net的价值不只是比U-Net高出的那零点几个百分点的Dice更在于它把残差连接引入医学分割领域后开启的一组新的结构设计思路。理解了块内残差、跨层级联跳、上下采样衔接这些核心逻辑后续无论迁移到哪种医学分割场景都能快速定位问题并给出合理的结构改进方案。把基础结构吃透远比盲目堆砌新的注意力模块和Transformer结构管用。