news 2026/9/23 15:08:36

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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 + ReLU,UNet 的基本单元""" def __init__(self, in_ch, out_ch): super().__init__() self.block = nn.Sequential( nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) def forward(self, x): return self.block(x)

padding=1保证卷积后空间尺寸不变,bias=False是因为后面接了 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, t=2): super().__init__() self.t = t self.conv = nn.Sequential( nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) def forward(self, x): for _ in range(self.t): x = self.conv(x) + x # 残差累加 return x

t=2是 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, bias=False), nn.BatchNorm2d(F_int), ) self.W_x = nn.Sequential( nn.Conv2d(F_l, F_int, 1, bias=False), nn.BatchNorm2d(F_int), ) self.psi = nn.Sequential( nn.Conv2d(F_int, 1, 1, bias=False), nn.BatchNorm2d(1), nn.Sigmoid(), ) self.relu = nn.ReLU(inplace=True) 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卷积在这里只做通道变换,不改变空间尺寸,所以gx的空间尺寸必须一致——如果解码器上采样后和编码器特征差一个像素,加和会直接报错,这是改网络时最常见的翻车点。

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, transform=None): 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(image=img, mask=mask) 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 系数和 IoU,Dice 对分割边界的敏感度比像素准确率高,像素准确率在极不平衡数据上能到 95% 却什么都没分出来。

import torch import torch.nn as nn 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 = probs.view(-1) targets = targets.view(-1) intersection = (probs * targets).sum() dice = (2. * intersection + self.smooth) / (probs.sum() + targets.sum() + self.smooth) return 1 - dice

smooth=1.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 或 SGD,UNet 系列对学习率不算特别敏感,但 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_ch=3, out_ch=1).to(device) criterion = torch.nn.BCEWithLogitsLoss() optimizer = torch.optim.Adam(model.parameters(), lr=1e-4) loader = DataLoader(dataset, batch_size=4, shuffle=True, num_workers=4) 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_ch=3对应 RGB 输入,灰度图改成 1;out_ch=1是二值分割,多类改成类别数并换 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 memory,batch 降到 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_dirmask_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 4

R2 系列学习率减半是因为残差累加让有效梯度变大,同样的 lr 下更容易震荡。跑完用 evaluation.py 统一评估,把结果整理成表:

模型DiceIoU参数量单 epoch 耗时
UNet基准基准最少最快
R2UNet略高略高中等中等
Attention-UNet边界更准略高中等中等
AttentionR2UNet通常最高最高最多最慢

这张表不是让你照抄数值,而是提醒你对比时把参数量和耗时一起记,否则容易陷入「精度高一点但慢三倍」的取舍困境。验证模型有没有真的学到东西,除了看指标,我习惯把预测掩码叠加到原图上存下来,肉眼看一下边界是不是贴合,指标高但边界糊的情况在医学图像里不少见。

从那以后我每次换数据集,都强制先跑一个 epoch 的小样本过拟合测试:拿 4 张图训练 200 步,看 loss 能不能降到接近 0。降不下去说明数据管道或标签有问题,别急着调模型。这个习惯帮我省了很多在错误方向上调参的时间。希望帮到你。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/23 15:07:12

形状记忆合金弹簧的MATLAB仿真:从相变原理到参数设计

简介:面向材料科学、力学与智能结构领域的研究者和工程师,这套仿真资源围绕形状记忆合金(SMA)的马氏体—奥氏体相变机理,提供可在MATLAB环境中直接运行的数值模拟程序。资源以仿真脚本为主线,综合考虑温度变…

作者头像 李华
网站建设 2026/9/23 14:58:50

正念接纳:心理学中的自我成长与情绪管理

1. 正念接纳的本质与价值上周在咖啡馆遇见老同学小林,她盯着咖啡杯苦笑:"每次照镜子都觉得这里不够好、那里要调整,连发朋友圈都要修图半小时。"这让我想起五年前那个不敢素颜见人的自己。如今能坦然接受眼角的细纹,并非…

作者头像 李华
网站建设 2026/9/23 14:56:55

脉冲S参数测量技术:原理、应用与实战解析

1. 脉冲S参数测量的核心价值与挑战在射频微波领域,脉冲S参数测量正成为评估高速数字电路、雷达系统和功率放大器动态特性的重要手段。与传统连续波测量相比,脉冲激励能更真实地模拟器件在实际工作中的瞬态响应。我曾在某相控阵雷达T/R组件测试中深有体会…

作者头像 李华
网站建设 2026/9/23 14:56:28

化疗药物激活STING通路重塑肿瘤免疫微环境

1. 研究背景与临床意义化疗药物与免疫系统的相互作用一直是肿瘤治疗领域的热点课题。吉西他滨和顺铂作为临床常用的标准化疗方案,在多种实体瘤(如非小细胞肺癌、膀胱癌等)中展现出明确的治疗效果。但传统观点认为,化疗主要通过直接…

作者头像 李华
网站建设 2026/9/23 14:54:25

农学论文真相✅别堆砌田间方案+作物生理数据硬凑综述

农学、作物栽培学、育种学、植物营养、耕作学方向本科生研究生狠狠共情! 农学文献综述,是农林类典型“栽培方案同质化、试验描述高度撞文”的查重重灾区! 综述高频覆盖:作物栽培调控、品种选育、施肥管理、抗旱抗逆生理、耕作模式…

作者头像 李华