医学图像分割这个方向,这两年卷得确实厉害。但不管你跑多少新模型,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.9
- PyTorch 2.1.0 + CUDA 11.8
- torchvision 0.16.0
- opencv-python 4.8.1
- albumentations 1.3.1
- numpy 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(p=0.5), A.VerticalFlip(p=0.5), A.RandomRotate90(p=0.5), A.Normalize(mean=0.5, std=0.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, transform=None): 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(image=image, mask=mask) 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和ReLU:
import torch import torch.nn as nn import torch.nn.functional as F class DoubleConv(nn.Module): def __init__(self, in_channels, out_channels, mid_channels=None): super().__init__() if mid_channels is None: mid_channels = out_channels self.double_conv = nn.Sequential( nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=1, bias=False), nn.BatchNorm2d(mid_channels), nn.ReLU(inplace=True), nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1, bias=False), nn.BatchNorm2d(out_channels), nn.ReLU(inplace=True), ) 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_size=3, padding=1, bias=False) self.bn1 = nn.BatchNorm2d(out_channels) self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1, bias=False) self.bn2 = nn.BatchNorm2d(out_channels) self.relu = nn.ReLU(inplace=True) self.shortcut = nn.Identity() if in_channels != out_channels: self.shortcut = nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False), 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有个实现细节值得注意:我在这里用的是bias=False。因为在卷积后面紧跟BatchNorm,BN会在计算时对输入做标准化和偏移,卷积层的bias会被标准化的过程抵消掉,留着反而浪费参数量,还容易引入冗余。这是从ResNet论文里沿用下来的做法,复现的时候建议保持一致。
3.2 编码器与解码器的实现
编码器部分的责任是做特征提取。我用一个Encoder类来管理:一个残差块负责特征提取,一个最大池化负责降低分辨率,这样循环堆叠4次。
class Encoder(nn.Module): def __init__(self, in_channels=3, 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_size=2, stride=2)) 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_size=2, stride=2) ) 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, size=skip.shape[2:], mode="bilinear", align_corners=True) x = torch.cat([x, skip], dim=1) 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_channels=3, out_channels=1, 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_size=1) 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_channels=3, out_channels=1, 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乘以预测和目标区域的交集,除以预测区域加目标区域的总和。加上一个平滑项防止分母为0:
class DiceLoss(nn.Module): def __init__(self, smooth=1.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(dim=1) union = probs_flat.sum(dim=1) + targets_flat.sum(dim=1) 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 Loss:
bce = 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, threshold=0.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()注意这里threshold=0.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_channels=3, out_channels=1).cuda() optimizer = AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4) scheduler = CosineAnnealingLR(optimizer, T_max=100, eta_min=1e-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(f"Epoch {epoch+1}/{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或BCE+Dice组合。如果是全白——也就是模型输出全部为正类,这通常发生在你的掩码标签里有大量区域被错误地置为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跑通,再去替换论文中的模块,这样哪个模块有效、哪个模块没用,一目了然。最后再提一句:复现网络结构本身只是开始,数据质量、损失函数和训练策略的分寸把握,才是决定分割效果的那道坎。希望这篇文章能帮你少踩几个坑,省下来的时间,拿去做实验本身吧。