U-Net到Attention R2U-Net:PyTorch实现四种医学图像分割模型对比
发布时间:2026/9/16 5:13:49 作者:尧图编辑部 阅读量:1,286

简介面向医学图像分割与深度学习入门进阶人群这份资源提供了基于PyTorch的U-Net、R2U-Net、Attention U-Net和Attention R2U-Net四种模型的完整实现并附带数据集与训练说明便于对比不同结构在同一任务上的表现。压缩包内共8个文件其中7个Python脚本分别承担网络定义、数据集加载、数据预处理、求解器配置、训练主程序与评估等模块另有1个Markdown说明文件梳理运行流程。整体包体仅12KB轻量精简适合动手实践。目前已有555人学习下载。通过学习这套源码可以掌握经典分割网络的搭建思路理解循环残差模块和注意力门控的改进逻辑同时借助配套数据完成端到端训练验证。代码按模块划分从数据加载、模型构建到训练评估形成完整闭环方便替换数据集进行二次开发是入门图像分割与模型对比实验的高性价比参考。1. 四个U-Net变体分割任务里到底该选哪个医学影像分割里U-Net是绝大多数团队的首选基线但很多人在实际项目里遇到的情况是基线能跑精度卡住改损失函数、调学习率都纹丝不动。这时候真正值得动手的方向不是换backbone而是改网络结构本身。R2U-Net在U-Net基础上引入循环残差卷积让每个卷积块在时间维度上复用参数小数据集上更稳Attention U-Net在跳跃连接处加注意力门控把解码器特征作为引导信号抑制背景区域响应。Attention R2U-Net把两者叠加是四者中表达力最强但训练也最重的。这套基于PyTorch的源码把四个模型放在同一个工程里network.py、solver.py、evaluation.py等模块划分清晰自带数据集和训练步骤说明适合做模型复现、对比实验或直接改到自己的分割项目里。2. U-Net到R2U-Net从跳跃连接到循环残差的结构演进2.1 U-Net的编码-解码骨架与跳跃连接U-Net的结构本质上是两条路径左侧编码器通过卷积加池化逐步降低空间分辨率、增加通道数右侧解码器通过上采样恢复分辨率中间靠跳跃连接把编码器每一层的特征直接拼到解码器对应层上。这里的关键点是concat而不是addconcat保留了通道维度上的独立信息让解码器既能拿到高层语义又能拿到浅层边界纹理。在network.py中U-Net的基础卷积块就是标准的双卷积结构# network.py 中 U-Net 的基础卷积块 def double_conv(in_channels, out_channels): return nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size3, padding1), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue), nn.Conv2d(out_channels, out_channels, kernel_size3, padding1), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue) )这段代码里kernel_size3配合padding1保证特征图尺寸在卷积前后不变这是U-Net能稳定做跳跃连接的前提。BatchNorm放在卷积和激活之间作用是稳定训练分布在分割任务里几乎必加否则深层网络很容易出现梯度震荡。两个卷积堆叠形成一个block编码器每个stage用它提取特征然后接2x2 maxpooling下采样。跳跃连接的价值在于经过四次下采样后解码器的特征图分辨率只有输入的1/16单纯靠上采样恢复出来的边缘是模糊的。把编码器对应层的特征拼过来相当于给解码器直接提供了高分辨率的细节信息这是U-Net在医学分割上比FCN强出明显一截的根本原因。2.2 R2U-Net的循环残差卷积模块R2U-Net改动的地方不在跳跃连接而在卷积块本身。它把双卷积替换成循环残差卷积块RRCNNBlock核心思想是同一个卷积核在时间步上重复使用。比如t2时一个3x3卷积对同一张特征图做两次卷积中间夹BN和ReLU权重是共享的。这样网络变深了但新增的参数量只来自BN卷积核还是那一组。# network.py 中 R2U-Net 的循环卷积模块 class RecurrentConv(nn.Module): def __init__(self, in_channels, out_channels, t2): super().__init__() self.t t self.conv nn.Conv2d(in_channels, out_channels, kernel_size3, padding1) self.bn nn.BatchNorm2d(out_channels) def forward(self, x): for i in range(self.t): x self.conv(x) x self.bn(x) x F.relu(x, inplaceTrue) return x class RRCNNBlock(nn.Module): def __init__(self, in_channels, out_channels, t2): super().__init__() self.rrcnn nn.Sequential( RecurrentConv(in_channels, out_channels, t), RecurrentConv(out_channels, out_channels, t) ) self.skip nn.Conv2d(in_channels, out_channels, kernel_size1) def forward(self, x): return self.skip(x) self.rrcnn(x)这里的t就是循环次数默认取2。第一次循环把in_channels映射到out_channels第二次循环在out_channels内部做特征细化。skip路径用1x1卷积把输入通道数对齐到out_channels残差相加时维度才一致。这种设计的好处是循环展开t次等价于t层共享权重的卷积理论上感受野随着t扩大但参数量几乎不增长。两个模型放在一起对比会更清楚对比项U-NetR2U-Net基础卷积块double_conv双卷积RRCNNBlock循环卷积残差循环时间步无t2可调跳跃连接编码器特征直接concat编码器特征直接concat参数量较低略高主要是BN和1x1卷积训练耗时快约为U-Net的1.3到1.6倍适用场景数据量大、追求速度小数据集、纹理复杂2.3 循环结构到底带来了什么R2U-Net在论文里的出发点是解决U-Net在数据量不足时特征提取不充分的问题。循环卷积让同一组参数在不同时间步上处理特征相当于给网络增加了一个隐式的深度维度在相同epoch下能比U-Net学到更细的纹理差异。尤其在病灶边界模糊、目标与背景灰度接近的影像上这种差异能被Dice分数直接体现出来。代价是显存占用和训练时间。循环卷积在反向传播时要展开t步显存开销接近直接加t层卷积所以实际使用时t不要超过3否则小显卡直接out of memory。如果你的数据集规模很大比如上万张自然图像U-Net的简单结构反而更稳R2U-Net的收益就不明显了。3. Attention U-Net与Attention R2U-Net注意力门控机制怎么改解码器3.1 注意力门控与SE、CBAM的本质区别很多人第一次接触Attention U-Net时都会误以为它像SE模块那样做通道注意力或者像CBAM那样通道加空间双分支。实际上Attention U-Net里的注意力门控Attention Gate是完全不同的思路它有两个输入一个是来自跳跃连接的编码器特征x另一个是来自解码器深层的门控信号g。g经过上采样后分辨率还是比x小但语义层次更深用它去调制x的空间权重让网络知道哪些位置才是真正需要关注的目标区域。这个思路和NLP里seq2seq模型在decoder侧使用attention的动机是一致的decoder每一步生成时回看encoder的不同位置并赋予不同权重。Attention U-Net把同样的逻辑搬到图像分割的跳跃连接上只不过权重从概率分布变成了二维空间注意力图。3.2 AttentionGate的PyTorch实现在network.py里AttentionGate的实现比较固定核心就是两个1x1卷积加一个sigmoid# network.py 中 Attention U-Net 的注意力门控模块 class AttentionGate(nn.Module): def __init__(self, in_channels, gating_channels, inter_channelsNone): super().__init__() inter_channels in_channels if inter_channels is None else inter_channels self.W_g nn.Conv2d(gating_channels, inter_channels, kernel_size1) self.W_x nn.Conv2d(in_channels, inter_channels, kernel_size1) self.psi nn.Conv2d(inter_channels, 1, kernel_size1) self.relu nn.ReLU(inplaceTrue) self.sigmoid nn.Sigmoid() def forward(self, x, g): # x: 跳跃连接传入的编码器特征 # g: 解码器上采样后的门控信号 g1 self.W_g(g) x1 self.W_x(x) out self.relu(g1 x1) out self.psi(out) attn self.sigmoid(out) return x * attn参数的对应关系要理解清楚in_channels是编码器特征x的通道数gating_channels是解码器特征g的通道数inter_channels是中间计算维度。W_g把g压缩到inter_channelsW_x把x也压缩到同样维度两者逐元素相加后过ReLU再经过psi压缩成单通道sigmoid映射到0到1之间。最后把注意力权重乘回x实现空间维度的特征重标定。这里用1x1卷积而不是3x3是刻意的注意力门控只需要做通道对齐和特征融合不需要引入额外的空间感受野。inter_channels如果取in_channels的一半可以明显减少计算量分割效果不会有太大损失。3.3 四种模型在network.py中的组合方式Attention R2U-Net不是重新设计一个网络而是把RRCNNBlock和AttentionGate拼在一起编码器和解码器的卷积块换成循环残差结构跳跃连接不再直接concat而是先过AttentionGate再加到解码器特征上。四种模型的关系可以整理成一张表模型编码器/解码器基础块跳跃连接方式U-Netdouble_conv直接concatR2U-NetRRCNNBlock直接concatAttention U-Netdouble_convAttentionGate加权后concatAttention R2U-NetRRCNNBlockAttentionGate加权后concat实际跑下来Attention R2U-Net在目标小、背景占比高的数据集上优势最明显因为注意力门控天然抑制了非目标区域的特征响应。但它的训练时间也是四个模型里最长的如果项目对推理速度有硬性要求不建议直接上这个。3.4 注意力图的可视化验证训练完成后把AttentionGate输出的attn权重图保存下来是最直观的验证手段# 可视化注意力权重图 import matplotlib.pyplot as plt import torch.nn.functional as F # attn 形状为 [B, 1, H, W]sigmoid 输出 attn_resized F.interpolate(attn, size(256, 256), modebilinear, align_cornersTrue) plt.imshow(attn_resized[0, 0].cpu().detach().numpy(), cmapjet) plt.axis(off) plt.savefig(attention_map.png, bbox_inchestight)这段代码里的attn是在forward里把AttentionGate中间层输出单独return出来的。如果注意力图上的高亮区域零散分布在背景边缘说明门控没学好常见原因是训练epoch不够或者学习率偏大。正常学到的注意力图应该集中在目标轮廓内部背景区域接近0。4. 从dataset.py到solver.py这套PyTorch工程怎么跑通训练全流程4.1 源码模块划分这套源码的组织方式和大部分PyTorch分割工程类似network.py放四个模型定义dataset.py负责读取原始图像和mask并做预处理data_loader.py把dataset包装成DataLoader并支持打乱和多进程加载solver.py封装训练循环、验证和模型保存逻辑main.py是命令行入口evaluation.py单独做评估指标计算misc.py放绘图、参数统计之类的工具函数。各模块职责单一改模型只动network.py改数据只动dataset.py调试起来很省事。4.2 dataset.py的核心代码与处理逻辑dataset.py里最关键的是图像和mask必须用完全相同的resize尺寸否则训练时模型学到的空间对应关系在推理时会错位。常见做法是统一缩放到256x256图像用双线性插值mask用最近邻插值# dataset.py 中分割数据集的核心实现 class SegDataset(Dataset): def __init__(self, image_dir, mask_dir, image_size256): self.image_dir image_dir self.mask_dir mask_dir self.image_size image_size self.images sorted(os.listdir(image_dir)) self.masks sorted(os.listdir(mask_dir)) assert len(self.images) len(self.masks) def __getitem__(self, idx): img_path os.path.join(self.image_dir, self.images[idx]) mask_path os.path.join(self.mask_dir, self.masks[idx]) image cv2.imread(img_path) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) image cv2.resize(image, (self.image_size, self.image_size)) image image.astype(np.float32) / 255.0 mask cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) mask cv2.resize(mask, (self.image_size, self.image_size), interpolationcv2.INTER_NEAREST) image torch.from_numpy(image).permute(2, 0, 1).float() mask torch.from_numpy(mask).long() return image, maskmask用INTER_NEAREST而不是默认的双线性插值这一点很关键。双线性插值会在类别边界处产生中间值比如0和1之间插出0.4这不是一个合法类别索引CrossEntropyLoss计算时会出错或者学出模糊边界。image转成torch张量后通过permute把HWC变成CHW这个顺序不能错PyTorch卷积层默认输入是通道在前。提示如果数据集的mask是彩色标注图需要先做颜色到类别索引的映射不能直接当灰度图读进来否则每个像素值都不是合法的类别id。4.3 solver.py训练循环与超参数设置solver.py里封装了optimizer、scheduler、loss和训练迭代。分割任务里最常用的组合是Adam加CrossEntropyLoss学习率初始1e-4配合ReduceLROnPlateau在验证集loss不再下降时自动衰减# solver.py 中训练核心配置 self.optimizer torch.optim.Adam(model.parameters(), lr1e-4) self.scheduler torch.optim.lr_scheduler.ReduceLROnPlateau( self.optimizer, modemin, factor0.1, patience5 ) self.criterion nn.CrossEntropyLoss() # 单个batch训练步骤 for images, masks in self.train_loader: images images.cuda() masks masks.cuda() outputs model(images) # 前向传播 loss self.criterion(outputs, masks) # 计算loss self.optimizer.zero_grad() # 梯度清零 loss.backward() # 反向传播 self.optimizer.step() # 参数更新zero_grad必须在backward之前调用因为PyTorch的梯度是累积的不清零会把上一个batch的梯度累加到当前batch上导致loss震荡甚至不收敛。三步的顺序不能换这是个写PyTorch训练循环最容易忽略的细节。超参数的参考配置可以按下面这张表来设定超参数常见取值设置原因image_size256显存和细节保留的折中batch_size812G显存下U-Net可以跑到8到16optimizerAdam分割任务收敛比SGD稳定lr1e-4预训练特征不需要过大的学习率schedulerReduceLROnPlateau验证集loss停滞时降lrepochs60-100R2U-Net需要更多轮次收敛4.4 main.py命令行入口与模型切换main.py负责解析命令行参数并启动训练。README里给出的典型执行方式类似python main.py --model attention_r2unet --epochs 80 --batch_size 8 --lr 1e-4--model参数用来切换四种模型取值一般是unet、r2unet、attention_unet、attention_r2unet四选一。如果数据集路径不是默认目录再加上--data_dir或者--train_dir参数指定。这套工程比自己在网上零散找单模型源码强的地方就在于结构统一、参数入口一致四个模型跑出来的指标可以直接横向对比不用为每个模型单独写一套训练脚本。4.5 训练过程中容易翻车的几个细节R2U-Net和Attention R2U-Net训练时显存占用会比U-Net高出一截batch_size要适当调小。模型定义里的t循环次数直接决定显存上限先确认一下有没有把t写大。另外eval和train之间的模式切换很容易漏训练时每次迭代前要model.train()验证和保存模型前要model.eval()漏掉的话BatchNorm的running_mean和running_var会持续更新推理结果可能出现诡异的偏移。5. evaluation.py里的验证细节指标计算与模型复现的常见坑5.1 Dice与IoU的实现细节evaluation.py里通常是逐张图计算Dice和IoU最后取平均。实现上有几个细节影响最终数值# evaluation.py 中 DICE 与 IoU 的计算实现 def dice_coef(pred, mask, smooth1e-5): pred (pred 0.5).float() intersection (pred * mask).sum() return (2.0 * intersection smooth) / (pred.sum() mask.sum() smooth) def iou_score(pred, mask, smooth1e-5): pred (pred 0.5).float() intersection (pred * mask).sum() union pred.sum() mask.sum() - intersection return (intersection smooth) / (union smooth)pred是模型输出经过sigmoid之后的概率图必须先阈值化成0/1再算指标。如果不做阈值化0.7和0.3的概率值直接参与计算Dice会被虚高这个指标就失去了对比意义。smooth参数的作用是防止分子分母同时为0尤其在目标特别小的图像上如果没有smooth整张图全是背景时会得到0/0的异常值。5.2 评估时最容易忽略的resize对齐问题训练时图像缩放到256x256评估阶段如果把原始尺寸直接送进模型输出和原始mask的空间对应关系是错的Dice必然偏低。正确做法是评估和训练保持完全一样的预处理流程模型输出后如果需要和原图尺寸的mask对比用F.interpolate把概率图上采样回原尺寸阈值化后再算指标。上采样的插值方式也要统一一般双线性即可但mask的resize始终用最近邻。5.3 用misc.py工具快速定位问题misc.py里一般存了count_parameters之类的工具函数训练前先打印模型参数量能很快发现网络结构是否正确加载。举个例子U-Net参数通常在1300万到3100万之间如果打印出来的参数量少了一个量级多半是编码器层数配错了。注意如果换数据集后loss能下降但Dice始终上不去先检查mask预处理看类别索引有没有从0开始、resize插值方式对不对这两个问题占了分割项目排错的大头。本文还有配套的精品资源点击获取