2026/9/23 15:08:43

PyTorch图像分割实战:UNet、R2UNet、Attention-UNet与AttentionR2UNet对比

PyTorch图像分割实战:UNet、R2UNet、Attention-UNet与AttentionR2UNet对比 简介这份资源面向图像分割方向的深度学习学习者与研究者提供基于Pytorch实现的UNet、R2UNet、Attention-UNet及AttentionR2UNet四种经典算法的完整项目代码可用于医学影像分析、自动驾驶、视频监控等场景下的分割任务实践与对比实验。压缩包共14个文件约257KB包含7个Python脚本网络结构、数据加载、训练与评估等核心模块、5张网络结构示意图、1个运行脚本和1份说明文档结构清晰便于快速复现与二次开发。目前已有239人学习下载。读者可借此掌握编码器-解码器架构、残差连接缓解梯度消失、注意力门控突出关键区域等关键设计并对照结构图理解各变体的差异适合作为课程设计、科研入门或算法选型的实战参考。1. 从一份能直接跑通的 PyTorch 图像分割工程说起医学影像里要把肝脏和肿瘤分开遥感图里要标出每一块耕地工业质检里要抠出划痕区域——这些任务的共同点是把每个像素归到某一类也就是图像分割。很多人第一次接触分割都是从 UNet 开始的结构对称、代码不长、论文好懂但真到自己动手往往卡在数据怎么读、损失怎么算、模型怎么改这几步上。这份基于 PyTorch 实现的工程把 UNet、R2UNet、Attention-UNet、AttentionR2UNet 四个变体放在同一套训练框架里网络定义、数据加载、训练循环、评估脚本都拆成了独立文件改一处就能对比不同结构的效果。它适合已经会写基础 PyTorch 训练循环、想系统对比分割模型改进思路的人也适合拿它当自己项目的骨架把 dataset.py 换成自己的数据就能跑起来。2. 四个网络变体的结构差异与选型逻辑2.1 UNet 的编码器-解码器与跳跃连接UNet 的核心是下采样提特征、上采样恢复分辨率再用跳跃连接把编码器的高分辨率特征拼到解码器对应层。编码器每经过一次池化空间尺寸减半、通道数翻倍感受野变大语义信息变强解码器做转置卷积或插值上采样把语义信息还原回原图尺寸。跳跃连接解决的是上采样过程中细节丢失的问题——池化丢掉的边缘、纹理通过拼接直接补回来。这也是 UNet 在医学图像上表现好的原因器官边界往往就是几个像素宽的灰度过渡没有跳跃连接这些边界在上采样后基本糊掉。工程里 network.py 把这一结构写成了可复用的模块。常见做法是把双层卷积抽成一个DoubleConv编码器和解码器各调用若干次这样改通道数或层数时只动一处。下面是我一般会用的写法和工程里的组织方式一致import torch import torch.nn as nn class DoubleConv(nn.Module): 两次 3x3 卷积 BN ReLUUNet 的基本单元 def __init__(self, in_ch, out_ch): super().__init__() self.block nn.Sequential( nn.Conv2d(in_ch, out_ch, kernel_size3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, kernel_size3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.block(x)padding1保证卷积后空间尺寸不变biasFalse是因为后面接了 BN偏置会被 BN 的减均值操作抵消省掉能少一点参数量。BN 放在卷积和 ReLU 之间是标准顺序训练时对 batch 统计量做归一化推理时用滑动平均这一点在评估脚本里要记得切model.eval()否则 BN 还在用当前 batch 的统计量单张图推理结果会飘。2.2 R2UNet 的循环残差块怎么加R2UNet 在 UNet 基础上把每个 DoubleConv 换成了循环残差块Recurrent Residual Block。残差连接让输入直接加到输出上缓解深层网络的梯度消失循环结构则是把同一个卷积单元在时间步上展开两次让特征反复 refinement。具体到实现一个 R2 块里有两个子块第一个子块的输出加上原始输入再送进第二个子块第二个子块的输出再加上第一个子块的输出。这样梯度可以沿着残差路径直接回传训练更深的网络时不容易出现前面几层几乎不更新情况。class RecurrentBlock(nn.Module): 循环卷积块同一组卷积在时间步上复用 t 次 def __init__(self, out_ch, t2): super().__init__() self.t t self.conv nn.Sequential( nn.Conv2d(out_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x): for _ in range(self.t): x self.conv(x) x # 残差累加 return xt2是 R2UNet 论文里的默认设置时间步再多收益递减、显存和耗时线性增长。这里把卷积和残差写在循环里等价于把同一个模块展开成两层但参数共享参数量比堆两层独立卷积少。要注意的是残差累加要求输入输出通道一致所以 R2 块一般放在通道数不变的位置通道变化的那一层还是用普通卷积过渡。2.3 Attention Gate 把注意力加在哪一层Attention-UNet 的关键是注意力门控Attention Gate它作用在跳跃连接上解码器当前层的特征作为门控信号编码器对应层的特征作为被筛选对象两者经过加和、激活后生成一个空间注意力图再乘回编码器特征。效果是让解码器在拼接时更关注和当前语义相关的区域抑制背景响应。医学图像里病灶往往只占一小块没有注意力时解码器会把大量背景纹理也拼进来注意力门控相当于给跳跃连接加了一道软掩码。class AttentionGate(nn.Module): 注意力门控g 为门控信号(解码器)x 为被筛选特征(编码器) def __init__(self, F_g, F_l, F_int): super().__init__() self.W_g nn.Sequential( nn.Conv2d(F_g, F_int, 1, biasFalse), nn.BatchNorm2d(F_int), ) self.W_x nn.Sequential( nn.Conv2d(F_l, F_int, 1, biasFalse), nn.BatchNorm2d(F_int), ) self.psi nn.Sequential( nn.Conv2d(F_int, 1, 1, biasFalse), nn.BatchNorm2d(1), nn.Sigmoid(), ) self.relu nn.ReLU(inplaceTrue) def forward(self, g, x): g1 self.W_g(g) x1 self.W_x(x) att self.relu(g1 x1) att self.psi(att) # 空间注意力图值域 0~1 return x * att # 加权后的编码器特征F_int是中间通道数一般取F_l // 2太大显存吃紧、太小注意力图分辨率不够。1x1卷积在这里只做通道变换不改变空间尺寸所以g和x的空间尺寸必须一致——如果解码器上采样后和编码器特征差一个像素加和会直接报错这是改网络时最常见的翻车点。2.4 AttentionR2UNet 的组合顺序AttentionR2UNet 就是把 R2 块和注意力门控拼在一起编码器和解码器的基本单元用 R2 块跳跃连接处插注意力门控。组合顺序上先做 R2 特征提取再在跳跃连接上做注意力加权最后拼接。不要反过来先注意力再 R2因为注意力图是基于当前特征生成的先加权再进 R2 块R2 的残差累加会把注意力权重稀释掉。工程里四个模型共用同一套训练和评估流程切换模型只需要改 main.py 里的模型名参数这也是它适合做对比实验的原因。3. 数据加载与训练流程的落地细节3.1 dataset.py 与 data_loader.py 的分工dataset.py 负责单样本的读取和预处理data_loader.py 负责批处理、打乱和多进程加载。分割任务的数据集和分类不一样输入是图像标签也是同尺寸的掩码所以__getitem__要同时返回 image 和 mask且两者必须做完全相同的几何变换翻转、旋转、裁剪否则图像和标签会错位。常见做法是把随机变换的参数先采样一次再分别作用到 image 和 mask 上。import os import numpy as np from torch.utils.data import Dataset from PIL import Image class SegmentationDataset(Dataset): def __init__(self, img_dir, mask_dir, transformNone): self.img_dir img_dir self.mask_dir mask_dir self.transform transform self.names sorted(os.listdir(img_dir)) def __len__(self): return len(self.names) def __getitem__(self, idx): name self.names[idx] img np.array(Image.open(os.path.join(self.img_dir, name)).convert(RGB)) mask np.array(Image.open(os.path.join(self.mask_dir, name)).convert(L)) mask (mask 127).astype(np.float32) # 二值化按实际阈值调整 if self.transform: augmented self.transform(imageimg, maskmask) img, mask augmented[image], augmented[mask] img img.transpose(2, 0, 1).astype(np.float32) / 255.0 return img, mask[None, ...] # mask 加通道维和输出对齐convert(L)把掩码转成单通道灰度mask 127是二值化阈值如果标签是多类这里要改成保留类别索引、不二值化。transpose(2,0,1)把 HWC 转成 CHW这是 PyTorch 卷积层的输入格式。mask 加[None, ...]是为了和模型输出的[B, 1, H, W]对齐少了这一维损失函数广播时可能不报错但算出来的值不对属于那种不翻车但结果玄学变差的坑。3.2 损失函数与评估指标的搭配二值分割常用 BCEWithLogitsLoss 或 Dice Loss前者对每个像素独立算交叉熵后者直接优化预测和标签的重叠度。类别极不平衡时比如病灶只占 5% 像素BCE 会被背景主导Dice 更稳。工程里 evaluation.py 一般会同时算 Dice 系数和 IoUDice 对分割边界的敏感度比像素准确率高像素准确率在极不平衡数据上能到 95% 却什么都没分出来。import torch import torch.nn as nn class DiceLoss(nn.Module): def __init__(self, smooth1.0): super().__init__() self.smooth smooth def forward(self, logits, targets): probs torch.sigmoid(logits) probs probs.view(-1) targets targets.view(-1) intersection (probs * targets).sum() dice (2. * intersection self.smooth) / (probs.sum() targets.sum() self.smooth) return 1 - dicesmooth1.0防止分母为零尤其是预测和标签都全零的样本。view(-1)把 batch 和空间维拉平Dice 是全局统计量逐样本算再平均和整体算会有差异训练时用整体算梯度更稳。实际训练里我一般把 BCE 和 Dice 按 1:1 加权BCE 提供稳定的逐像素梯度Dice 负责拉高重叠度单用 Dice 在训练初期梯度噪声偏大。3.3 main.py 里的训练循环与超参main.py 串起数据、模型、损失和优化器。分割任务的 batch size 受显存限制通常比分类小4 到 8 是常见起点配合num_workers开多进程读数据。学习率用 1e-3 到 1e-4配 Adam 或 SGDUNet 系列对学习率不算特别敏感但 R2 结构因为残差累加梯度幅值偏大学习率要适当调小。import torch from torch.utils.data import DataLoader from network import UNet, R2UNet, AttentionUNet, AttentionR2UNet device torch.device(cuda if torch.cuda.is_available() else cpu) model AttentionUNet(in_ch3, out_ch1).to(device) criterion torch.nn.BCEWithLogitsLoss() optimizer torch.optim.Adam(model.parameters(), lr1e-4) loader DataLoader(dataset, batch_size4, shuffleTrue, num_workers4) for epoch in range(100): model.train() for img, mask in loader: img, mask img.to(device), mask.to(device) optimizer.zero_grad() out model(img) loss criterion(out, mask) loss.backward() optimizer.step()in_ch3对应 RGB 输入灰度图改成 1out_ch1是二值分割多类改成类别数并换 CrossEntropyLoss。zero_grad()必须在backward()前否则梯度会跨 batch 累加。训练完一个 epoch 记得在验证集上跑model.eval()加torch.no_grad()评估脚本 evaluation.py 就是干这个的它加载权重、算指标、存预测图改数据路径时两个文件都要同步改。4. 环境搭建与训练排查的常见坑4.1 CUDA 版本和 PyTorch 对不上现象是torch.cuda.is_available()返回 False或者 import torch 时报找不到 cudart 相关动态库。原因是 pip 装的 PyTorch 默认是 CPU 版或者 CUDA 版本和驱动不匹配。解决是先nvidia-smi看驱动支持的 CUDA 上限再去 PyTorch 官网按对应版本选安装命令不要直接pip install torch。装完用python -c import torch; print(torch.version.cuda)确认编译时的 CUDA 版本和驱动版本是两回事。4.2 图像和掩码尺寸不一致导致拼接报错现象是训练到跳跃连接拼接时size mismatch或者上采样后尺寸差一个像素。原因是输入图像尺寸不是 16 的整数倍UNet 下采样四次每次除二奇数尺寸会在某层出现向下取整上采样时对不回去。解决是把输入统一 resize 到 16 的倍数如 256、512或者在拼接前用F.interpolate把解码器特征对齐到编码器尺寸。我一般直接在 dataset 里 resize省得网络里到处判断。4.3 损失不下降但也不报错现象是 loss 在 0.69 附近震荡Dice 接近 0。原因是标签没二值化或者 mask 的通道维没加对导致损失函数把背景全预测成 0 也能拿到低 loss。排查方法是取一个 batch 打印mask.unique()和mask.shape确认标签只有 0 和 1、形状是[B,1,H,W]。另一个常见原因是学习率太大R2 结构下 1e-3 容易发散降到 1e-4 再看。4.4 显存溢出但 batch size 已经很小现象是CUDA out of memorybatch 降到 1 还报。原因是注意力门控和 R2 块都会额外占显存尤其是注意力图在每层跳跃连接都生成一份。解决是先用 UNet 跑通流程再换 R2 或 Attention 版本或者用torch.cuda.empty_cache()清理缓存把num_workers调小避免多进程各自占显存。混合精度训练torch.cuda.amp能省一半左右显存但要注意 Dice Loss 在 fp16 下的数值稳定性一般 loss 计算留在 fp32。4.5 评估指标和训练指标对不上现象是训练时 loss 降得很好evaluation.py 跑出来 Dice 却很低。原因是评估时忘了model.eval()BN 还在用 batch 统计量或者评估用的预处理和训练不一致比如训练做了归一化评估没做。排查时把评估脚本里的预处理单独打印出来和 dataset 对比确认model.eval()和torch.no_grad()都加了。5. 换自己的数据集与模型对比的实操技巧拿到这份工程后最直接的用法是先把四个模型在同一份数据上跑一遍看 Dice 和 IoU 的差距。我一般会固定随机种子、固定数据划分只改 main.py 里的模型名其他超参不动这样对比才有意义。换自己数据时dataset.py 里改img_dir和mask_dir确认掩码是单通道、像素值是 0/1 或 0/255多类任务把out_ch改成类别数、损失换成 CrossEntropyLoss、评估指标按类算再平均。# 固定种子跑对比实验四个模型依次训练 python main.py --model unet --epochs 100 --lr 1e-4 --batch_size 4 python main.py --model r2unet --epochs 100 --lr 5e-5 --batch_size 4 python main.py --model attunet --epochs 100 --lr 1e-4 --batch_size 4 python main.py --model attr2unet --epochs 100 --lr 5e-5 --batch_size 4R2 系列学习率减半是因为残差累加让有效梯度变大同样的 lr 下更容易震荡。跑完用 evaluation.py 统一评估把结果整理成表模型DiceIoU参数量单 epoch 耗时UNet基准基准最少最快R2UNet略高略高中等中等Attention-UNet边界更准略高中等中等AttentionR2UNet通常最高最高最多最慢这张表不是让你照抄数值而是提醒你对比时把参数量和耗时一起记否则容易陷入「精度高一点但慢三倍」的取舍困境。验证模型有没有真的学到东西除了看指标我习惯把预测掩码叠加到原图上存下来肉眼看一下边界是不是贴合指标高但边界糊的情况在医学图像里不少见。从那以后我每次换数据集都强制先跑一个 epoch 的小样本过拟合测试拿 4 张图训练 200 步看 loss 能不能降到接近 0。降不下去说明数据管道或标签有问题别急着调模型。这个习惯帮我省了很多在错误方向上调参的时间。希望帮到你。本文还有配套的精品资源点击获取