Res U-Net实战:PyTorch实现医学图像分割的完整指南

📅 发布时间:2026/10/4 5:57:00
Res U-Net实战:PyTorch实现医学图像分割的完整指南
医学图像分割这个方向这两年卷得确实厉害。但不管你跑多少新模型U-Net这条线始终绕不开。今天这篇不写花活就老老实实把Res U-Net这个经典改版在PyTorch里怎么复现、怎么调参、怎么避坑完整过一遍。我会从网络结构拆解讲到损失函数再到训练推理全流程最后把实际跑数据时遇到的那些报错和诡异现象也一并整理出来。无论你是刚入坑分割任务的学生还是想快速在业务里验证一个baseline的工程师这篇文章都能直接当参考手册用。先说结论Res U-Net本质上就是在U-Net的每个卷积块里加上了残差连接Residual Connection这个改动看起来不大但带来的训练稳定性和精度收益是实打实的。尤其是医学图像通常样本少、噪声多残差结构能让网络在加深的同时不掉点这对分割任务来说是救命级的特性。1. Res U-Net的设计思路残差连接为什么能在分割网络里站稳脚跟1.1 从U-Net到Res U-Net它到底改了什么经典的U-Net是2015年提出的编码器-解码器结构编码器负责逐层下采样提取语义特征解码器负责逐层上采样恢复分辨率。中间通过跳跃连接把编码器每一层的特征图拼接到解码器对应层这样边界信息和语义信息能互补。它靠这个设计霸榜医学分割很多年这不用多说。Res U-Net的改进点很直接把原来那个双3x3卷积ReLU的基础块换成了带残差连接的卷积块。也就是说每个stage的输入除了进入卷积层堆叠之外还会通过一个shortcut跳过路径在输出端和卷积结果逐元素相加。如果输入输出通道数不一致就在shortcut上补一个1x1卷积来对齐维度。这个改动的意义可以用一句话概括让网络退而求其次变得容易——如果某个stage的卷积学不到有效特征残差连接允许网络直接把输入传递下去不会因为无效特征提取而丢信息。这在实际训练中意味着网络可以适当加深而不用担心梯度消失和退化问题。1.2 为什么残差特别适合医学图像分割医学图像和自然图像有个很重要的差异样本量小而且目标区域往往只占整幅图的极小比例。CT里一个几毫米的结节MRI里一小块病灶放在512x512的图像里可能就几十个像素。这就导致两个问题第一训练数据少网络容易过拟合深层网络更容易把噪声当作特征记住。残差结构相当于在网络中加入了恒等映射的先验等于告诉网络你学不到新东西的时候就保持原样这天然是一种隐式正则化。第二目标太小梯度信号弱。尤其是Dice Loss这类区域重叠型损失函数当预测和真实标签完全不相交时梯度会变得非常不稳定。残差路径给梯度提供了一条高速公路让backward的梯度信号即使穿过很多层也能传导回来。我在实验里的感受是同样epoch数下Res U-Net的收敛速度明显快于普通U-Net验证集Dice通常能高出2到3个点。1.3 Res U-Net的宏观结构速览从整体上看Res U-Net的结构是编码器4个stage每个stage包含一个残差块和一个2x2最大池化通道数依次从64升到128、256、512瓶颈层最底层再用一个残差块把512通道扩到1024这一层没有池化解码器4个stage每个stage先做上采样转置卷积或双线性插值然后与编码器对应层的特征在通道维度拼接再经过一个残差块最终输出一个1x1卷积把通道数降到类别数激活函数根据损失函数选择这里有个细节要注意编码器的空间分辨率会逐层减半解码器对应层的分辨率也要逐层恢复一半。在PyTorch里做拼接前必须确保两边的尺寸完全一致否则torch.cat会直接报错。实际中因为padding和卷积步长的设置有些层会出现偶数的尺寸偏差后面我会专门讲怎么处理这类问题。2. 复现前的准备环境搭建与数据组织2.1 环境搭建与依赖版本PyTorch的安装没什么难度如果只用CPU跑小数据集直接装CPU版本就够了有GPU的话就装对应CUDA版本的。这里列一个我常用的版本组合供参考Python 3.9PyTorch 2.1.0 CUDA 11.8torchvision 0.16.0opencv-python 4.8.1albumentations 1.3.1numpy 1.24.4装好之后用一段最简代码验证CUDA是否可用import torch print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0) if torch.cuda.is_available() else CPU Only)如果输出True和显卡型号说明环境没问题。这一步值得花两分钟确认因为后面所有训练的报错有一半都能追溯到CUDA和PyTorch版本不匹配。2.2 数据集的目录组织与预处理医学图像分割任务的数据集通常长这样一个文件夹放原始图像jpg/png/nii.gz一个文件夹放对应的掩码标注。我建议大家在复现的时候统一用images和masks两个文件夹来管理文件重名后缀不同这样写Dataset类最省事。预处理里有几个关键抉择图像尺寸医学图像原始分辨率通常很大要提前统一resize到固定尺寸。我常用256x256来做前期实验512x512做最终训练。尺寸太大对显存和训练速度影响很大起步阶段用256x256足够验证模型有没有问题。归一化医学图像的灰度范围很不统一通常是16位整型范围从0到4095或者更高。我习惯直接用最大值做Min-Max归一化到0-1区间比ImageNet的均值和标准差归一化更适合医学数据。数据增强分割任务里最常用的增强是随机水平翻转、随机垂直翻转和随机旋转90度。这个组合对医学图像非常友好因为解剖结构的朝向虽然固定但轻微的几何扰动不会改变语义信息反而能显著提升泛化能力。注意做翻转和旋转时图像和掩码必须使用完全相同的随机种子否则标签就错位了。如果用albumentations库它会自动帮你处理这个对齐问题所以我建议直接用albumentations。import albumentations as A from albumentations.pytorch import ToTensorV2 train_transform A.Compose([ A.Resize(256, 256), A.HorizontalFlip(p0.5), A.VerticalFlip(p0.5), A.RandomRotate90(p0.5), A.Normalize(mean0.5, std0.5), ToTensorV2() ])2.3 Dataset类的实现细节写Dataset类的时候有几个坑值得注意。一个是掩码的质量标注文件通常会有抗锯齿边缘像素值可能不是严格的0和255读进来之后做个二值化把大于127的值统一置为255小于等于127的置为0这样后面计算Dice时才不会出现中间值导致的偏差。另一个是掩码的通道维度PyTorch模型输入的shape是(N, C, H, W)所以掩码也要保留通道维度shape是(N, 1, H, W)。下面是一个可以直接复制的Dataset实现import os import cv2 import torch from torch.utils.data import Dataset class SegmentationDataset(Dataset): def __init__(self, images_dir, masks_dir, transformNone): self.images_dir images_dir self.masks_dir masks_dir self.images sorted(os.listdir(images_dir)) self.transform transform def __len__(self): return len(self.images) def __getitem__(self, idx): image_name self.images[idx] image_path os.path.join(self.images_dir, image_name) mask_path os.path.join(self.masks_dir, image_name.replace(.png, _mask.png)) image cv2.imread(image_path, cv2.IMREAD_COLOR) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) mask cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) mask (mask 127).astype(uint8) * 255 if self.transform: augmented self.transform(imageimage, maskmask) image augmented[image] mask augmented[mask] mask (mask 0.5).float().unsqueeze(0) return image, mask掩码这边把像素范围直接归一化到0和1的浮点数后续损失函数里省得反复转换。3. Res U-Net的核心代码实现这一节是整篇文章的重头戏我会把网络每一层都拆开来讲并附上完整的PyTorch实现代码。你可以直接照着敲也可以理解之后按照自己的需求改。3.1 基础卷积块与残差块的实现先在代码层面定义两个基础模块普通双卷积块DoubleConv和残差块ResidualBlock。普通双卷积块就是U-Net里的老配方两个3x3卷积每个卷积后面接BatchNorm和ReLUimport torch import torch.nn as nn import torch.nn.functional as F class DoubleConv(nn.Module): def __init__(self, in_channels, out_channels, mid_channelsNone): super().__init__() if mid_channels is None: mid_channels out_channels self.double_conv nn.Sequential( nn.Conv2d(in_channels, mid_channels, kernel_size3, padding1, biasFalse), nn.BatchNorm2d(mid_channels), nn.ReLU(inplaceTrue), nn.Conv2d(mid_channels, out_channels, kernel_size3, padding1, biasFalse), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.double_conv(x)Res U-Net里这个基础块被升级成残差版本。残差块的forward流程是输入x先进两个卷积块然后和自身相加最后过ReLU。如果输入输出通道数不一致需要在shortcut上补一个1x1卷积class ResBlock(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.conv1 nn.Conv2d(in_channels, out_channels, kernel_size3, padding1, biasFalse) self.bn1 nn.BatchNorm2d(out_channels) self.conv2 nn.Conv2d(out_channels, out_channels, kernel_size3, padding1, biasFalse) self.bn2 nn.BatchNorm2d(out_channels) self.relu nn.ReLU(inplaceTrue) self.shortcut nn.Identity() if in_channels ! out_channels: self.shortcut nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size1, 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 out identity out self.relu(out) return out有个实现细节值得注意我在这里用的是biasFalse。因为在卷积后面紧跟BatchNormBN会在计算时对输入做标准化和偏移卷积层的bias会被标准化的过程抵消掉留着反而浪费参数量还容易引入冗余。这是从ResNet论文里沿用下来的做法复现的时候建议保持一致。3.2 编码器与解码器的实现编码器部分的责任是做特征提取。我用一个Encoder类来管理一个残差块负责特征提取一个最大池化负责降低分辨率这样循环堆叠4次。class Encoder(nn.Module): def __init__(self, in_channels3, features(64, 128, 256, 512)): super().__init__() self.blocks nn.ModuleList() self.pools nn.ModuleList() for feature in features: self.blocks.append(ResBlock(in_channels, feature)) self.pools.append(nn.MaxPool2d(kernel_size2, stride2)) in_channels feature def forward(self, x): skip_features [] for block, pool in zip(self.blocks, self.pools): x block(x) skip_features.append(x) x pool(x) return x, skip_features解码器的实现比编码器多两步先上采样然后拼接跳跃连接传来的特征再过残差块。上采样有两种常见选择转置卷积和双线性插值。转置卷积是U-Net原版的风格参数可学习但偶尔会引入棋盘格伪影双线性插值没有可学习参数但结果平滑棋盘格问题几乎不存在。我在医学图像分割里更倾向于用转置卷积因为医学图像对边界细节敏感可学习的上采样核更能适应不同器官的形状。如果你发现输出有奇怪的格子纹路再换成双线性插值也来得及。class Decoder(nn.Module): def __init__(self, features(512, 256, 128, 64)): super().__init__() self.upconvs nn.ModuleList() self.blocks nn.ModuleList() for idx, feature in enumerate(features): self.upconvs.append( nn.ConvTranspose2d(feature * 2, feature, kernel_size2, stride2) ) self.blocks.append(ResBlock(feature * 2, feature)) def forward(self, x, skip_features): skip_features skip_features[::-1] for upconv, block, skip in zip(self.upconvs, self.blocks, skip_features): x upconv(x) if x.shape ! skip.shape: x F.interpolate(x, sizeskip.shape[2:], modebilinear, align_cornersTrue) x torch.cat([x, skip], dim1) x block(x) return x这里我必须专门强调一个很多新手踩过的坑跳跃连接拼接前上采样结果和skip特征图的尺寸必须完全一致。由于编码器里有两次下采样如果原始图尺寸不是4的倍数拼接时就很可能出现差一个像素的尴尬局面。我上面在代码里加了一段通用的处理逻辑先用F.interpolate把x强制缩放到和skip一样的尺寸再接torch.cat。虽然多了一步运算但换来了尺寸的完全兼容在实际训练中能省掉大量调试时间。3.3 完整网络组装最后把编码器、瓶颈层bottleneck、解码器拼起来再在末尾加一个1x1卷积输出指定类别的分割图class ResUNet(nn.Module): def __init__(self, in_channels3, out_channels1, features(64, 128, 256, 512)): super().__init__() self.encoder Encoder(in_channels, features) self.bottleneck ResBlock(features[-1], features[-1] * 2) self.decoder Decoder(features[::-1]) self.final_conv nn.Conv2d(features[0], out_channels, kernel_size1) def forward(self, x): x, skip_features self.encoder(x) x self.bottleneck(x) x self.decoder(x, skip_features) logits self.final_conv(x) return logits为什么瓶颈层要扩通道到原来的两倍因为经过4次下采样之后特征图的分辨率已经变成原来的1/16空间信息丢失很严重此时需要用更大的通道数来保留更丰富的抽象特征。这也是U-Net系列一贯的设计逻辑越靠后的特征图空间分辨率越低通道数就要越高。用一段简单代码来验证网络能否正常前向传播model ResUNet(in_channels3, out_channels1, features(64, 128, 256, 512)) fake_input torch.randn(2, 3, 256, 256) output model(fake_input) print(output.shape) # 期望输出 torch.Size([2, 1, 256, 256])如果输出shape是(2, 1, 256, 256)网络结构就是对的。这里建议你实际跑一遍用这种方式快速确认网络搭建没有尺寸问题。4. 损失函数与评估指标4.1 分割任务里怎么选损失函数医学图像分割的场景下类别不平衡是常态。一个512x512的图像里器官或者病灶可能只占几十到几百像素常规的CrossEntropy Loss会让模型轻松陷入全预测背景的局部最优。所以分割任务最常用的损失是Dice Loss它直接优化区域重叠度对像素级别的类别不平衡不太敏感。Dice Loss的公式不复杂2乘以预测和目标区域的交集除以预测区域加目标区域的总和。加上一个平滑项防止分母为0class DiceLoss(nn.Module): def __init__(self, smooth1.0): super().__init__() self.smooth smooth def forward(self, logits, targets): probs torch.sigmoid(logits) probs_flat probs.view(probs.size(0), -1) targets_flat targets.view(targets.size(0), -1) intersection (probs_flat * targets_flat).sum(dim1) union probs_flat.sum(dim1) targets_flat.sum(dim1) dice (2.0 * intersection self.smooth) / (union self.smooth) return 1.0 - dice.mean()实际使用中我建议用Dice Loss和BCE Loss的加权组合。只靠Dice Loss对单个像素梯度不够友好加上BCE可以稳定整个训练过程。常见的配比是0.5倍的Dice Loss加0.5倍的BCE Lossbce nn.BCEWithLogitsLoss() dice DiceLoss() logits model(image) loss 0.5 * bce(logits, mask) 0.5 * dice(logits, mask)这个组合是我在多个数据集上试过之后比较稳的方案比单纯Dice Loss收敛快比单纯BCE Loss的精度高。4.2 评估指标Dice系数和IoU训练过程中的评测指标最常用的就是Dice系数和IoU。这两个指标都很好理解Dice相当于预测和真实区域的重叠程度加权平均IoU直接算交集除以并集。下面是一个简单的验证集评测函数def calculate_dice(pred_mask, true_mask, threshold0.5): pred_mask torch.sigmoid(pred_mask) pred_mask (pred_mask threshold).float() intersection (pred_mask * true_mask).sum() dice (2.0 * intersection) / (pred_mask.sum() true_mask.sum() 1e-7) return dice.item()注意这里threshold0.5只是一个初始值。在验证集上你完全可以搜索一个最优阈值用0.3到0.7中间步长0.05去遍历取验证Dice最高的那个阈值作为推理时的阈值。这个技巧虽然简单但往往能白捡0.5到1个点的Dice。5. 训练与推理的完整流程5.1 训练循环的写法训练代码本身不复杂但有几个经验值得分享。第一优化器用AdamW比Adam在权值衰减上更规范在分割任务里效果也更好。学习率可以先用1e-4作为基准如果loss下降太慢再作调整。第二加一个余弦退火学习率调度器能避免到了训练后期还在用大步长震荡。from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR model ResUNet(in_channels3, out_channels1).cuda() optimizer AdamW(model.parameters(), lr1e-4, weight_decay1e-4) scheduler CosineAnnealingLR(optimizer, T_max100, eta_min1e-6)完整训练循环大致如下for epoch in range(epochs): model.train() train_loss 0.0 for images, masks in train_loader: images, masks images.cuda(), masks.cuda() logits model(images) loss 0.5 * bce(logits, masks) 0.5 * dice(logits, masks) optimizer.zero_grad() loss.backward() optimizer.step() train_loss loss.item() * images.size(0) scheduler.step() avg_loss train_loss / len(train_loader.dataset) print(fEpoch {epoch1}/{epochs}, Loss: {avg_loss:.4f})关于model.train()和model.eval()的切换再啰嗦一句BN层在训练和推理时的行为完全不同。训练时用当前batch的均值和方差推理时用running mean和running variance。如果不小心在推理时没有切到eval模式BN层的统计量会一直用最后一个batch的结果会非常离谱尤其在batch size小的时候更明显。5.2 预测与后处理预测阶段同样要切成model.eval()并且用torch.no_grad()关闭梯度追踪省显存也提升速度model.eval() with torch.no_grad(): logits model(image_batch) probs torch.sigmoid(logits) pred_mask (probs threshold).float()医学图像分割的后处理里最常用的一招就是去除小的连通域。因为模型很容易在远离目标的位置输出一些零星的小噪声块面积只有几个到十几个像素用cv2.connectedComponentsWithStats扫一遍把面积小于某个阈值的连通域直接剔除就能得到干净很多的分割结果。这个后处理步骤不影响整体Dice但视觉上会好看非常多也能提升一些基于区域的评估指标。6. 训练过程中的常见问题与排查技巧6.1 显存不足与图片尺寸的取舍训练医学图像分割模型最常遇到的就是CUDA out of memory。我在最初复现的时候也踩过这个坑。512x512输入配合64初始通道的Res U-Net在8GB显存上batch size设到8基本就爆炸了。解决办法有几个方向第一个降低输入分辨率从512降到256显存直接变成原来的四分之一第二个减小batch size比如从8降到2或4这通常是最快的办法第三个关闭混合精度训练以外的冗余内存也就是在模型中加torch.cuda.empty_cache()虽然治标不治本修但能顶住后续突发峰值。如果你想保持512分辨率又想提升速度可以考虑用自动混合精度训练PyTorch原生提供了支持一般能省一半显存训练速度还能提升一截。6.2 预测结果全是黑的或全白的排查训练了一段时间之后发现验证集上模型输出的掩码全是背景全黑或者全是前景全白这种情况我从经验来看大概率是损失函数出了问题。尤其是只使用BCE时类别不平衡会让模型倾向于把全部像素都预测为背景因为背景占的比例太大损失下降的主力方向就是预测得更像背景。解决方法是立刻切换到Dice Loss或BCEDice组合。如果是全白——也就是模型输出全部为正类这通常发生在你的掩码标签里有大量区域被错误地置为255建议检查数据加载和预处理环节里二值化阈值设置是否正确。6.3 验证Dice不错但预测效果肉眼很差的场景还有一种情况比较迷惑验证集上的Dice很漂亮但你随便拿一张图出来看分割结果边界糙得没法看或者出现了很多细碎的小孔洞。这通常是因为模型对边界像素的划分不够精确而Dice系数对边界误差并不敏感——毕竟只差几个像素的偏移对区域重叠率来说就是小数点后两三位的事。这时候能做的有两个方向一是把输入分辨率提上去从256提到512边界细节一般能有可感知的提升二是加一些边缘感知的损失项比如在交叉熵损失里给边界像素更高的权重或者直接用带边界惩罚的损失变体。不过如果你只是做baseline复现前一个方向就够用了。6.4 训练Loss震荡不下降的排查Loss在训练初期不降反升或者一直在高位震荡这种情况第一个要怀疑的就是学习率设置过大。医学图像任务的数据分布和自然图像差异很大初始学习率用1e-4比2e-3稳妥得多。第二个要怀疑的就是数据归一化不一致图像可能被归一化到了-1到1而掩码是0到1模型输入输出的尺度不同也会让训练很难收敛统一归一化到0-1区间最省心。另外如果发现训练Dice上去了但验证Dice迟迟不涨大概率是过拟合了。这时候除了早停还可以把数据增强开大点不只是翻转旋转再加点随机亮度对比度扰动和弹性形变对医学图像是安全且有效的增强手段。7. 一点实操体会Res U-Net这个网络严格说不是那种能给你带来SOTA新突破的模型但它作为医学图像分割的baseline价值极高——结构清晰、实现不难、训练稳定几乎所有熟人圈的进阶模型都是在这个框架之上做的改进。在我实际跑过的肾部CT、眼底血管等几个数据集上Res U-Net的收敛速度和最终Dice都比普通U-Net要更稳。如果你正在复现一篇新论文里的分割网络又担心它训练起来不稳定我的建议是先用Res U-Net把baseline跑通再去替换论文中的模块这样哪个模块有效、哪个模块没用一目了然。最后再提一句复现网络结构本身只是开始数据质量、损失函数和训练策略的分寸把握才是决定分割效果的那道坎。希望这篇文章能帮你少踩几个坑省下来的时间拿去做实验本身吧。