
做了几年医学图像分割每次带新人复现网络最常遇到的不是模型写不出来而是写出来之后要么不收敛、要么收敛了但指标不对劲。Res U-Net就是这样光看结构图会觉得很简单U-Net加残差块而已可真在PyTorch里一行行敲下来通道数对不上、BatchNorm在batch_size很小的时候发神经、skip connection拼接出错哪哪都是问题。这篇我把自己完整复现Res U-Net的过程分享出来从它为什么能解决医学影像分割的问题开始到每个模块的PyTorch实现、数据管道、训练技巧和评估避坑适合正在入门医学图像分割、又不想只调现成高级接口的同学。1. 为什么医学影像分割任务需要Res U-Net这种结构1.1 U-Net在医学分割里的统治地位与瓶颈医学图像分割跟自然图像语义分割不一样大部分场景下样本量少、标注成本极高而且不同组织、器官之间的边界常常是模糊的。早期FCN、SegNet这些网络在自然图像分割上表现不错但到了医学影像上就经常在边界处崩掉。U-Net能成为医学分割的默认基线核心原因是它的编码器-解码器结构加上skip connection让网络同时具备语义信息和空间细节。但U-Net有一个很实际的问题当编码器不断下采样网络层数加深之后训练早期梯度很难反传到浅层导致特征提取的不够充分。虽然U-Net本身没有像后来的ResNet那么深但在小样本医学任务里一旦初始化稍差或者BatchNorm更新不稳定训练就会很痛苦。另一个问题是U-Net的每个卷积子模块是普通卷积-卷积结构它没有一个“保底”通路让信息直接流过网络必须通过非线性变换才能传递特征这实际上限制了特征复用的效率。1.2 残差连接解决的核心矛盾加深网络与梯度消失Res U-Net的思路非常直接把U-Net里每一个普通的卷积子模块替换成残差子模块也就是在两层卷积之外加一条恒等捷径。恒等捷径带来的好处是反向传播时梯度可以从深层直接“抄近道”回到浅层即使主干卷积层的梯度很小网络依然能学到东西。这里有一个容易忽略的点残差块不是ELUs、不是Dropout它是一种对信息流的约束——让每一层在输入特征的“基础”上去学习一个残差增量。医学图像的特点决定了任务更适合这种结构。拿肝脏CT分割举例肝脏在相邻切片上的外观变化往往是渐进的残差连接能让隐藏层更容易学到“相对上一个特征图的微小变化”而不必重新编码整个器官外观。也就是说Res U-Net不仅解决了梯度消失还在特征表达能力上更契合医学图像的局部一致性。1.3 Res U-Net家族里你大概率会遇到的变体复现之前要意识到Res U-Net在文献里并不是一个唯一确定的结构。最早的Res U-Net是Zheng等在2018年提出的核心就是残差单元替换U-Net的卷积单元。后来又有很多变体比如编码器直接换成了ResNet34/ResNet50保留U-Net的跳跃结构也有在编码器用残差块、解码器用普通卷积块的。这些变体在论文里都叫Res U-Net复现时的代码差别很大评估结果也会不同。我自己复现时偏向采用最经典、最可控的结构编码器和解码器都用残留块下采样用步长为2的卷积上采样用转置卷积。这样结构非常统一用PyTorch写起来最顺也方便做消融实验。下面讲到的所有代码都基于这个经典结构。2. Res U-Net核心模块拆解从普通卷积块到残差卷积块2.1 残差卷积块的PyTorch实现先写一个最基础的残差块。它包含两个3×3卷积每个卷积后接BatchNorm和ReLU然后加上shortcut。当输入输出通道数一致时shortcut直接是恒等映射通道数不一致时用1×1卷积调整通道。import torch import torch.nn as nn class ResidualConv(nn.Module): def __init__(self, in_channels, out_channels, stride1): super(ResidualConv, 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) 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) ) else: self.shortcut nn.Identity() def forward(self, x): residual 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 residual return self.relu(out)这里我踩过一个坑inplaceTrue在ReLU里看起来省内存但如果你在做一些需要保留前向特征做对比的实验或者用torch.jit脚本化它会带来莫名其妙的问题。复现调参阶段建议先保留inplaceFalse跑通了再优化。2.2 编码器、瓶颈与解码器的结构计算Res U-Net主体可以理解为四个编码器层、一个瓶颈、四个解码器层每层通道数依次增加再对称减少。典型配置如下层输入通道输出通道操作Encoder Block 0164ResidualConvstride1Encoder Block 164128ResidualConvstride2Encoder Block 2128256ResidualConvstride2Encoder Block 3256512ResidualConvstride2Bottleneck5121024ResidualConvstride2Decoder Block 3512512512Upsample ResidualConvDecoder Block 2256256256Upsample ResidualConvDecoder Block 1128128128Upsample ResidualConvDecoder Block 0646464Upsample ResidualConv这里需要解释一下为什么解码器拼接后是512512。因为编码器第3层输出的特征图是512通道瓶颈输出的特征图上采样后也是512通道两者拼接后是1024通道再通过残余卷积压缩回512通道。很多新人在这里写代码报错通常是拼接维度写错或者在跳跃连接时拿错了特征图。2.3 上采样方式选择转置卷积还是插值解码器里上采样我用了转置卷积因为转置卷积可以和后面的残差卷积融合让网络在上采样过程中自己学一点可训练的参数。也有人用双线性插值然后接普通卷积。这两者在医学分割任务上都常见但转置卷积更容易出现棋盘伪影插值方式更稳。我的建议是如果图特别大、显存有限用F.interpolate(modebilinear)加上普通残差块更省显存。如果追求端到端可训练且不care显存转置卷积也行。复现阶段先用插值快速验证模型能不能收敛再用转置卷积优化指标。两者在代码上是三行以内的区别。3. PyTorch复现环境准备与数据管道细节3.1 环境配置与版本选择复现Res U-Net不需要最新的PyTorch但版本太老会遇到API差异问题。推荐Python 3.8以上PyTorch 1.10到2.x都行。我用的是1.13配合CUDA 11.7主要理由是稳定各种第三方库都兼容。创建环境建议直接用condaconda create -n resunet python3.9 conda activate resunet pip install torch1.13.1 torchvision0.14.1 --index-url https://download.pytorch.org/whl/cu117注意torchvision不是必须的但数据集加载、图像变换通常用到建议一起装。还要装上nibabel因为很多医学公开数据集是NIfTI格式不是普通png。3.2 处理图像与掩膜的Dataset类医学图像分割里最容易出问题的就是mask的读取。很多数据集标注类别的像素值是0、1、2但有些数据集是0和255甚至还有RGB三通道的color mask。写Dataset类时我习惯在读取后立刻归一化并断言mask里有且只有目标类别数。import numpy as np import cv2 from torch.utils.data import Dataset class MedicalImageDataset(Dataset): def __init__(self, image_paths, mask_paths, target_size(256, 256), transformNone): self.image_paths image_paths self.mask_paths mask_paths self.target_size target_size self.transform transform def __len__(self): return len(self.image_paths) def __getitem__(self, idx): image cv2.imread(self.image_paths[idx], cv2.IMREAD_GRAYSCALE) mask cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE) image cv2.resize(image, self.target_size, interpolationcv2.INTER_LINEAR) mask cv2.resize(mask, self.target_size, interpolationcv2.INTER_NEAREST) image image.astype(np.float32) / 255.0 mask mask.astype(np.float32) mask[mask 0] 1.0 image image[None, ...] # CHW mask mask[None, ...] if self.transform is not None: image, mask self.transform(image, mask) return torch.from_numpy(image), torch.from_numpy(mask)重点在mask的resize插值方式。图像我用INTER_LINEAR没问题但mask必须用INTER_NEAREST用线性插值会把0/1边界变成中间值网络评估时阈值化就废了。这是复现过程中最容易被忽略、也最影响指标的地方。3.3 数据增强的尺度拿捏医学图像数据增强不能太猛。翻转、随机旋转、小幅平移这些都会让网络更鲁棒。但不要为了刷指标而对mask做复杂的弹性形变除非你的数据集特别小且已经验证了这种形变符合真实场景。比如病变区域可能会因为扫描角度产生形变这时弹性形变是合理的。CT、MRI影像因为像素值范围大最好做z-score归一化不要简单除以255。一个可复现的增强方案随机水平翻转概率0.5随机旋转10度随机加减亮度最后统一resize到256×256。这套组合在大多数公开数据集上都能稳定涨点也不会引入额外的依赖。4. 网络主体逐段实现编码器、瓶颈与解码器4.1 编码器与跳跃连接的正确存储方式主体代码里最核心的地方在于保存每一层编码器输出供解码器拼接。有点粗糙的写法是用一个list存简单但容易搞混顺序。我建议用tuple返回并且给每层取清晰的名字。class ResUNet(nn.Module): def __init__(self, in_channels1, out_channels1, features[64, 128, 256, 512]): super(ResUNet, self).__init__() self.encoder1 ResidualConv(in_channels, features[0]) self.encoder2 ResidualConv(features[0], features[1], stride2) self.encoder3 ResidualConv(features[1], features[2], stride2) self.encoder4 ResidualConv(features[2], features[3], stride2) self.bottleneck ResidualConv(features[3], features[3] * 2, stride2) self.up4 nn.ConvTranspose2d(features[3] * 2, features[3], kernel_size2, stride2) self.decoder4 ResidualConv(features[3] features[3], features[3]) self.up3 nn.ConvTranspose2d(features[3], features[2], kernel_size2, stride2) self.decoder3 ResidualConv(features[2] features[2], features[2]) self.up2 nn.ConvTranspose2d(features[2], features[1], kernel_size2, stride2) self.decoder2 ResidualConv(features[1] features[1], features[1]) self.up1 nn.ConvTranspose2d(features[1], features[0], kernel_size2, stride2) self.decoder1 ResidualConv(features[0] features[0], features[0]) self.output nn.Conv2d(features[0], out_channels, kernel_size1)forward里要特别注意拼接顺序。有人习惯把编码器特征拼在前面解码器特征拼在后面模型照样能训但最好统一写作torch.cat((x, skip), dim1)这样代码可读性好不会在后续改动里埋雷。4.2 输出层与激活函数选择医学分割一般是二分类或多分类。二分类时输出层用1个通道加Sigmoid多分类用N个通道加Softmax。很多代码库里喜欢在模型里直接加Sigmoid但我更建议模型只输出logits在训练或评估时自行选择激活函数。这样可以用不一样的损失函数比如BCEWithLogitsLoss数值上更稳定。如果用DiceLoss也可以在logits上手动加Sigmoid效果更容易控制。4.3 前向传播的shape核对复现时最容易出错的环节就是forward里各层shape。简单起见可以在模型后面加一个测试函数def test_shape(): net ResUNet(in_channels1, out_channels1) x torch.randn((2, 1, 256, 256)) y net(x) print(y.shape)输入256×256时输出应该也是(2, 1, 256, 256)。如果你改了输入尺寸需要注意编码器下采样次数不能太多否则到解码器时空间分辨率太小信息损失严重。四层下采样对128×128的输入仍然可行但对64×64输入就已经很勉强了。一般训练都固定到256×256这个尺寸均衡了显存和精度。5. 训练策略与收敛性观察损失函数、学习率与显存控制5.1 为什么我推荐BCE加Dice的组合损失医学分割里正负样本极不平衡纯BCE容易让网络倾向于预测背景。纯Dice Loss又容易出现训练初期梯度不稳定。组合起来用是常见做法class BCEDiceLoss(nn.Module): def __init__(self, bce_weight0.5, dice_weight0.5): super(BCEDiceLoss, self).__init__() self.bce nn.BCEWithLogitsLoss() self.bce_weight bce_weight self.dice_weight dice_weight def forward(self, logits, targets): probs torch.sigmoid(logits) smooth 1e-5 bce_loss self.bce(logits, targets) intersection (probs * targets).sum() dice_loss 1 - (2.0 * intersection smooth) / (probs.sum() targets.sum() smooth) return self.bce_weight * bce_loss self.dice_weight * dice_loss训练前期bce_weight可以稍微调大一点帮助网络快速把背景区域学出来后期dice_weight大一点专注提升前景区域的重合度。但总体来说0.5/0.5已经能覆盖很多场景不需要过度调参。5.2 学习率、优化器和训练循环我常用AdamW初始学习率1e-4配合poly学习率衰减也就是def poly_lr(epoch, max_epochs, initial_lr1e-4, power0.9): return initial_lr * (1 - epoch / max_epochs) ** power也可以直接用torch.optim.lr_scheduler.LambdaLR。相比step衰减poly衰减在医学分割里更稳因为模型在后半段能缓慢收敛不会突然掉点。训练循环里除了记录loss至少要记录Dice系数。我习惯每个epoch结束在验证集上算一次Dice和IoU只看loss容易被数据不平衡骗了。Dice能直观反映分割效果尤其是肿瘤区域很小的时候loss可能降低得很缓慢Dice的变化更有意义。5.3 显存爆掉与batch_size困局医学图像比较大3D数据尤其夸张。如果输入是256×256的2D切片batch_size设为16一般没问题。如果用的是3D数据比如以块为单位输入batch_size往往只能设到2甚至1这时BatchNorm基本失效因为统计量抖动大。解决办法有两个。一是用梯度累积模拟更大batch_size而不是在实际维度上加大batch。二是把BatchNorm换成GroupNorm或者用torch.nn.SyncBatchNorm进行多卡同步。单卡情况下我觉得GroupNorm更省心。我没法把代码全部贴出来但核心就是nn.GroupNorm(num_groups16, num_channelschannels)在残差块里把BatchNorm替换掉就行这改动很小却在检测和医学分割任务中经常救急。5.4 训练过程中不能只看曲线还要看预测图等到训练到第20轮左右我把验证集上的原图、mask和预测结果拼在一起保存成一张图肉眼检查。这个习惯救了我很多次。比如有时候曲线很漂亮Dice到了0.9但预测图上的边缘完全黏在一起或者小块病灶全部漏掉。只靠指标是发现不了这些问题的。验证可视化可以很简陋三张图横着拼接保存成jpg。重点是原图、标签、预测三个图必须来自同一个样本而且预测要经过阈值0.5的处理不能直接存概率图。6. 评估、可视化与易遇到的坑从Dice系数到边界度量6.1 离线评估指标选择与正确实现评估时最忌讳只用一个Dice系数因为在病灶较小的情况下Dice会偏高或偏低不稳定。我一般同时报告Dice、IoU和Hausdorff距离。Dice与IoU计算方式相近Dice能反映重合度Hausdorff能反映边界最大偏差。尤其对医学图像来说边界最大误差比平均误差更让医生关注。def dice_coefficient(pred, target, smooth1e-5): pred (pred 0.5).float() intersection (pred * target).sum() return (2.0 * intersection smooth) / (pred.sum() target.sum() smooth) def iou_score(pred, target, smooth1e-5): pred (pred 0.5).float() intersection (pred * target).sum() union pred.sum() target.sum() - intersection return (intersection smooth) / (union smooth)注意计算Dice时要统一输入尺寸。如果数据管道里已经resize到256×256那评估也在256×256空间做。如果还要和原始mask尺寸比较需要把预测结果先resize回原始尺寸再用原始目标计算不能拿256×256上的结果当作最终效果因为插值本身会带来误差。6.2 复现一致性随机种子与确定性模式论文复现最容易碰到的麻烦是每次跑出来的结果不太一样。如果你想让结果尽可能可复现至少固定三处import random import numpy as np import torch def set_seed(seed42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic True torch.backends.cudnn.benchmark False但注意cudnn.benchmarkFalse会增加一些训练时间。在实际复现时我会先把它关掉以获得更稳定的结果等完全确定方案后再考虑打开benchmark加速。6.3 推理时碰到黑图或者全白图推理阶段最容易遇到的问题有三个。一是mask保存时用了cv2.imwrite但传入的是0~1的float32数组结果输出全黑。正确做法是先乘255再转uint8或者直接用PIL保存。二是预测时忘了对输入做和训练一样的归一化导致模型输入分布不一致输出一团糟。这个问题常见于读取新数据时没有除以255或者没有做z-score归一化。三是上采样到原图尺寸后再二值化。如果先二值化再做resize边缘会出现锯齿和空洞。正确做法是先插值概率图再阈值化这样边缘更平滑更接近真实分割掩膜。6.4 从训练到部署的常规打包方式训练结束后把模型权重、配置参数、预处理参数都整理好。权重没必要保存每个epoch我通常保留验证集Dice最高的那一个和最后一个用两个文件区分。推理脚本里最好只依赖模型结构定义和权重文件不依赖训练时的优化器状态。整个项目落到最后其实就是三件事数据管道正确、模型结构shape正确、训练策略稳定。Res U-Net的复现难点不在结构本身而在于任何一个环节的隐性错误都会在指标上体现得很延迟。我在实际使用中最受益的一个习惯是每修改一个环节先用一个很小的数据集跑两三个epoch确认loss能降下来再继续调参数这个方法帮我过滤掉了太多无效实验。如果你也想在自己的数据集上复现建议先从公开的小数据集跑通全流程再换到自己的数据。第一次跑通的意义不是拿到高分而是验证你的代码、数据、评估链路都是通的。后面调优只是时间问题。