ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

UNet、R2UNet、Attention-UNet 选哪个?PyTorch 实战对比与避坑指南

UNet、R2UNet、Attention-UNet 选哪个?PyTorch 实战对比与避坑指南 简介这份资源面向计算机视觉方向的学习者与研究者聚焦图像分割这一核心任务提供基于Pytorch实现的UNet、R2UNet、Attention-UNet与AttentionR2UNet四种经典网络结构的完整项目实战代码。内容覆盖医学影像分析、自动驾驶、视频监控等典型应用场景适合具备一定深度学习基础、希望深入理解分割算法演进脉络的中高级读者。资源包共14个文件包含7个Python脚本、5张网络结构示意图、1个运行脚本与1份说明文档压缩包约257KB代码与图示配合便于对照理解各模型的架构差异。目前已有239人学习下载。读者可从中获得可直接运行的训练与评估流程、数据集加载与网络定义模块以及四种变体在残差连接与注意力门控上的实现细节有助于快速复现实验并在此基础上开展改进与对比研究。1. 三个 UNet 变体摆在一起先搞清楚你该复现哪一个医学图像分割、广告牌图像分割系统、遥感地块提取这些场景里 UNet 几乎是默认起点。但真到动手时很多人会卡在同一个问题上UNet、R2UNet、Attention-UNet 到底选哪个我见过太多人一上来就冲 Attention-UNet结果训练半天 Dice 不升反降回头发现自己的数据集只有几百张图注意力模块根本没东西可学。这三个网络不是替代关系而是针对不同痛点的递进方案。UNet 解决的是“多尺度特征怎么融合”R2UNet 解决的是“深了以后梯度还传不传得动”Attention-UNet 解决的是“哪些区域值得花算力”。这篇笔记就按这个逻辑从环境搭建到三个模型跑通、再到训练自己的数据集把每一步的命令、参数和翻车点都写清楚。适合已经会 PyTorch 基础、想拿分割项目练手或落地的人。2. 环境搭建与数据准备别让 CUDA 版本成为第一个拦路虎2.1 PyTorch 环境搭建的版本对齐逻辑环境这一步翻车的人最多核心原因就一个PyTorch、CUDA、显卡驱动三者的版本没有对齐。常见做法是用 Anaconda 建独立环境避免和系统 Python 打架。先确认显卡驱动支持的 CUDA 上限再选 PyTorch 版本顺序不能反。# 查看显卡驱动和可支持的 CUDA 版本上限 nvidia-smi # 创建独立环境Python 版本建议 3.9 或 3.10 conda create -n unet_seg python3.10 -y conda activate unet_seg # 安装 PyTorchcu118 对应 CUDA 11.8按 nvidia-smi 结果调整 # 这条命令来自 PyTorch 官方索引不要用 pip install torch 裸装 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 # 验证 GPU 是否可用 python -c import torch; print(torch.__version__, torch.cuda.is_available())逻辑说明nvidia-smi右上角显示的 CUDA Version 是驱动能支持的最高版本不是已安装版本。PyTorch 自带的 CUDA 运行时只要不超过这个上限就能用。参数上cu118可以换成cu121但不要选超过驱动上限的。如果torch.cuda.is_available()返回 False九成是装成了 CPU 版用pip list | grep torch看版本号后面有没有cu后缀。提示没有独立显卡也能跑通代码把安装命令里的cu118换成cpu即可只是训练会慢到怀疑人生建议先用小图验证流程。2.2 分割数据集的目录结构与标签格式UNet 系列对数据格式不挑但标签必须是单通道的掩码图像素值就是类别索引。常见做法是整理成 images 和 masks 两个平行目录文件名一一对应。dataset/ ├── train/ │ ├── images/ # 原图jpg 或 png │ └── masks/ # 标签png单通道像素值 0/1/2... └── val/ ├── images/ └── masks/逻辑说明masks 用 png 而不是 jpg因为 jpg 有损压缩会污染像素值把类别 1 变成 0.98训练直接出玄学问题。二分类任务掩码像素值用 0 和 1多分类用 0 到 N-1。如果原始标签是彩色掩码比如标注工具导出的 RGB需要先转成单通道索引图这一步不做后面损失函数算出来的值全是错的。import numpy as np from PIL import Image # 把 RGB 彩色掩码转成单通道索引图 def rgb_mask_to_index(mask_path, color_map): mask np.array(Image.open(mask_path).convert(RGB)) index np.zeros(mask.shape[:2], dtypenp.uint8) for idx, color in enumerate(color_map): # 逐通道比对命中则该像素赋类别索引 match np.all(mask color, axis-1) index[match] idx return index # color_map 按你的标注颜色顺序填例如背景黑、目标白 color_map [(0, 0, 0), (255, 255, 255)] idx_mask rgb_mask_to_index(train/masks/001.png, color_map) Image.fromarray(idx_mask).save(train/masks/001_index.png)参数说明color_map的顺序决定了类别索引必须和训练时num_classes的语义一致。转换后建议抽查几张用np.unique看像素值是否只有预期的几个数。2.3 数据增强与 DataLoader 的最小实现分割任务的数据增强必须同时作用于原图和掩码且掩码只能用最近邻插值否则类别边界会被插出中间值。我一般用 albumentations它把图像和掩码的同步变换封装好了。import albumentations as A from albumentations.pytorch import ToTensorV2 # 训练增强几何变换对图和掩码同步生效 train_transform A.Compose([ A.Resize(256, 256), A.HorizontalFlip(p0.5), A.RandomRotate90(p0.5), A.Normalize(mean(0.485, 0.456, 0.406), std(0.229, 0.224, 0.225)), ToTensorV2(), ]) # 验证集只做 resize 和归一化不做随机变换 val_transform A.Compose([ A.Resize(256, 256), A.Normalize(mean(0.485, 0.456, 0.406), std(0.229, 0.224, 0.225)), ToTensorV2(), ])逻辑说明Normalize用的均值方差是 ImageNet 统计值因为骨干网络通常加载预训练权重输入分布要匹配。掩码经过ToTensorV2后是长整型张量配合CrossEntropyLoss使用。如果显存吃紧把Resize降到 128但小目标分割精度会掉这个取舍要自己权衡。3. UNet 基线先把最朴素的编码器解码器跑通3.1 UNet 的跳跃连接到底在解决什么UNet 的结构一句话概括编码器逐层下采样提语义解码器逐层上采样恢复分辨率跳跃连接把编码器同层的高分辨率特征直接拼到解码器。没有跳跃连接上采样出来的边界是糊的因为深层特征虽然知道“这是肿瘤”但不知道“边界在哪”。跳跃连接补的就是这个空间细节。这也是为什么 UNet 在小数据集上依然能打——它把浅层细节保留到了输出端。3.2 用 PyTorch 写一个可复用的 UNetimport torch import torch.nn as nn # 双层卷积块UNet 的基本单元 class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.conv(x) class UNet(nn.Module): def __init__(self, in_ch3, num_classes2, base64): super().__init__() # 编码器每次下采样通道翻倍 self.enc1 DoubleConv(in_ch, base) self.enc2 DoubleConv(base, base * 2) self.enc3 DoubleConv(base * 2, base * 4) self.enc4 DoubleConv(base * 4, base * 8) self.pool nn.MaxPool2d(2) # 瓶颈层 self.bottleneck DoubleConv(base * 8, base * 16) # 解码器上采样后与编码器特征拼接 self.up4 nn.ConvTranspose2d(base * 16, base * 8, 2, stride2) self.dec4 DoubleConv(base * 16, base * 8) self.up3 nn.ConvTranspose2d(base * 8, base * 4, 2, stride2) self.dec3 DoubleConv(base * 8, base * 4) self.up2 nn.ConvTranspose2d(base * 4, base * 2, 2, stride2) self.dec2 DoubleConv(base * 4, base * 2) self.up1 nn.ConvTranspose2d(base * 2, base, 2, stride2) self.dec1 DoubleConv(base * 2, base) self.out nn.Conv2d(base, num_classes, 1) def forward(self, x): e1 self.enc1(x) e2 self.enc2(self.pool(e1)) e3 self.enc3(self.pool(e2)) e4 self.enc4(self.pool(e3)) b self.bottleneck(self.pool(e4)) # 拼接时通道数翻倍所以 dec 的输入是 up 输出加对应编码器输出 d4 self.dec4(torch.cat([self.up4(b), e4], dim1)) d3 self.dec3(torch.cat([self.up3(d4), e3], dim1)) d2 self.dec2(torch.cat([self.up2(d3), e2], dim1)) d1 self.dec1(torch.cat([self.up1(d2), e1], dim1)) return self.out(d1)逻辑说明base控制首层通道数显存不够就降到 32 或 16。torch.cat沿通道维拼接所以解码器卷积的输入通道是上采样输出加编码器输出之和。ConvTranspose2d的stride2实现两倍上采样也可以用nn.Upsample加普通卷积替代前者有可学习参数后者更省显存。参数说明in_ch按输入图通道填RGB 是 3灰度是 1。num_classes是类别数二分类填 2不要填 1否则CrossEntropyLoss会报错。3.3 训练循环与 Dice 指标的落地写法import torch from torch.utils.data import DataLoader device torch.device(cuda if torch.cuda.is_available() else cpu) model UNet(in_ch3, num_classes2).to(device) optimizer torch.optim.Adam(model.parameters(), lr1e-3) criterion torch.nn.CrossEntropyLoss() def dice_coeff(pred, target, eps1e-6): # pred 是 argmax 后的类别索引target 是标签 pred pred.float() target target.float() inter (pred * target).sum() return (2 * inter eps) / (pred.sum() target.sum() eps) for epoch in range(50): model.train() for img, mask in train_loader: img, mask img.to(device), mask.to(device) optimizer.zero_grad() out model(img) loss criterion(out, mask) loss.backward() optimizer.step() # 验证阶段算 Dice model.eval() with torch.no_grad(): for img, mask in val_loader: img, mask img.to(device), mask.to(device) pred model(img).argmax(dim1) d dice_coeff(pred, mask)逻辑说明CrossEntropyLoss内部带 softmax所以模型输出直接是 logits不要在外面再加 softmax。Dice 计算前要argmax把 logits 转成类别索引。学习率 1e-3 是 Adam 的常用起点如果 loss 震荡就降到 1e-4。参数说明batch_size从 4 或 8 起步看显存。训练轮数不是越多越好分割任务通常 50 到 100 轮就能看出趋势过拟合了加早停。4. R2UNet把残差和循环卷积叠上去深网络才训得动4.1 残差循环块为什么能缓解梯度消失R2UNet 的核心是 Recurrent Residual Convolutional Unit简称 R2U。它做了两件事一是把残差连接塞进卷积单元让梯度有捷径可走二是让卷积单元循环复用同一组权重跑 T 次在不增加参数量的前提下加深有效感受野。普通 UNet 堆到很深时浅层梯度会衰减到接近零R2U 的残差路径把这个问题绕过去了。代价是训练变慢因为循环那部分要串行计算。4.2 R2U 单元与 R2UNet 的实现class R2U(nn.Module): def __init__(self, in_ch, out_ch, t2): super().__init__() self.t t # 循环次数 self.conv nn.Conv2d(in_ch, out_ch, 3, padding1) def forward(self, x): # 第一次循环输入是原始 x out torch.relu(self.conv(x)) for _ in range(self.t - 1): # 后续循环复用同一卷积输入叠加原始 x形成残差 out torch.relu(self.conv(x out)) return out class R2UNet(nn.Module): def __init__(self, in_ch3, num_classes2, base64, t2): super().__init__() self.enc1 R2U(in_ch, base, t) self.enc2 R2U(base, base * 2, t) self.enc3 R2U(base * 2, base * 4, t) self.pool nn.MaxPool2d(2) self.bottleneck R2U(base * 4, base * 8, t) self.up3 nn.ConvTranspose2d(base * 8, base * 4, 2, stride2) self.dec3 R2U(base * 8, base * 4, t) self.up2 nn.ConvTranspose2d(base * 4, base * 2, 2, stride2) self.dec2 R2U(base * 4, base * 2, t) self.up1 nn.ConvTranspose2d(base * 2, base, 2, stride2) self.dec1 R2U(base * 2, base, t) self.out nn.Conv2d(base, num_classes, 1) def forward(self, x): e1 self.enc1(x) e2 self.enc2(self.pool(e1)) e3 self.enc3(self.pool(e2)) b self.bottleneck(self.pool(e3)) d3 self.dec3(torch.cat([self.up3(b), e3], dim1)) d2 self.dec2(torch.cat([self.up2(d3), e2], dim1)) d1 self.dec1(torch.cat([self.up1(d2), e1], dim1)) return self.out(d1)逻辑说明R2U里x out是残差叠加self.conv在循环中共享权重。t越大感受野越广但训练时间线性增长一般取 2 或 3。注意这里为了代码清晰做了简化实际 R2U 的残差是加在卷积输出上再激活和标准实现略有出入但训练效果方向一致。参数说明t2是论文常用值显存不够先降base再降t。R2UNet 参数量比同base的 UNet 小因为循环共享权重但计算量更大。4.3 R2UNet 训练时的学习率与显存调优R2UNet 因为循环结构反向传播的图更深显存占用比 UNet 高。我一般把base从 64 降到 32batch_size保持 4学习率用 5e-4 而不是 1e-3因为残差结构对学习率更敏感太大容易震荡。如果训练时 loss 出现 NaN先检查是不是学习率过高再检查数据里有没有全黑的掩码导致 Dice 分母为零。# R2UNet 训练配置示例 model R2UNet(in_ch3, num_classes2, base32, t2).to(device) optimizer torch.optim.Adam(model.parameters(), lr5e-4) # 加梯度裁剪防止循环结构梯度爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)逻辑说明clip_grad_norm_把梯度范数限制在 1.0循环结构容易累积梯度裁剪是后悔药。参数max_norm从 1.0 起调太小会拖慢收敛。5. Attention-UNet注意力门控加在哪一层效果差很多5.1 注意力门控的机制与放置位置Attention-UNet 在跳跃连接上加了注意力门控用解码器当前层的特征作为门控信号去筛选编码器传过来的特征把无关区域的响应压下去。医学图像里病灶只占一小块背景占大头注意力门控能让网络聚焦到病灶。关键在于加的位置——加在浅层跳跃连接上筛的是边缘细节加在深层筛的是语义区域。常见做法是每一层跳跃连接都加但如果数据集小浅层加注意力反而会丢掉细节我一般只在后两层加。5.2 注意力门控模块与 Attention-UNet 实现class AttentionGate(nn.Module): def __init__(self, gate_ch, skip_ch, inter_ch): super().__init__() # 门控信号和跳跃特征分别做 1x1 卷积降到中间维度 self.W_g nn.Sequential( nn.Conv2d(gate_ch, inter_ch, 1), nn.BatchNorm2d(inter_ch), ) self.W_x nn.Sequential( nn.Conv2d(skip_ch, inter_ch, 1), nn.BatchNorm2d(inter_ch), ) # 合成后输出单通道注意力系数 self.psi nn.Sequential( nn.Conv2d(inter_ch, 1, 1), nn.BatchNorm2d(1), nn.Sigmoid(), ) self.relu nn.ReLU(inplaceTrue) def forward(self, g, x): # g 是解码器门控信号x 是编码器跳跃特征 g1 self.W_g(g) x1 self.W_x(x) # 尺寸不一致时对齐到 x 的尺寸 if g1.shape[2:] ! x1.shape[2:]: g1 nn.functional.interpolate(g1, sizex1.shape[2:], modebilinear, align_cornersFalse) att self.psi(self.relu(g1 x1)) return x * att # 逐像素加权逻辑说明W_g和W_x把两个输入降到同一中间维度再相加psi输出 0 到 1 的注意力图最后乘回跳跃特征。inter_ch一般取skip_ch的一半。插值对齐是因为上采样后的门控信号尺寸可能和跳跃特征差一两个像素不对齐会报维度错误。参数说明gate_ch是解码器特征的通道数skip_ch是编码器对应层的通道数填错会在torch.cat或相加时报错。5.3 把注意力门控接进 UNet 的跳跃连接class AttentionUNet(nn.Module): def __init__(self, in_ch3, num_classes2, base64): super().__init__() self.enc1 DoubleConv(in_ch, base) self.enc2 DoubleConv(base, base * 2) self.enc3 DoubleConv(base * 2, base * 4) self.pool nn.MaxPool2d(2) self.bottleneck DoubleConv(base * 4, base * 8) self.up3 nn.ConvTranspose2d(base * 8, base * 4, 2, stride2) # 只在后两层跳跃连接加注意力 self.att3 AttentionGate(base * 4, base * 4, base * 2) self.dec3 DoubleConv(base * 8, base * 4) self.up2 nn.ConvTranspose2d(base * 4, base * 2, 2, stride2) self.att2 AttentionGate(base * 2, base * 2, base) self.dec2 DoubleConv(base * 4, base * 2) self.up1 nn.ConvTranspose2d(base * 2, base, 2, stride2) self.dec1 DoubleConv(base * 2, base) self.out nn.Conv2d(base, num_classes, 1) def forward(self, x): e1 self.enc1(x) e2 self.enc2(self.pool(e1)) e3 self.enc3(self.pool(e2)) b self.bottleneck(self.pool(e3)) u3 self.up3(b) # 门控信号用上采样结果跳跃特征用编码器输出 a3 self.att3(u3, e3) d3 self.dec3(torch.cat([u3, a3], dim1)) u2 self.up2(d3) a2 self.att2(u2, e2) d2 self.dec2(torch.cat([u2, a2], dim1)) d1 self.dec1(torch.cat([self.up1(d2), e1], dim1)) return self.out(d1)逻辑说明注意力门控的输出a3和上采样结果u3拼接后送进解码器卷积。浅层e1不加注意力保留边缘细节。如果数据集目标占比很小可以三层都加但要注意浅层注意力可能让边界变糊。参数说明AttentionGate的第一个参数是门控信号通道第二个是跳跃特征通道这里两者相等是因为上采样输出和编码器同层通道数一致。如果改了base这些通道数要同步改。6. 避坑与排查三个模型训练时最容易翻车的五个地方6.1 现象loss 一直不降Dice 在 0.1 附近晃原因最常见的是标签像素值不是从 0 开始的连续整数比如背景是 255、目标是 0CrossEntropyLoss直接算错。其次是输入没归一化像素值 0 到 255 直接进网络梯度爆炸。解决用np.unique检查掩码像素值确保是 0 到num_classes-1。归一化用 ImageNet 均值方差或者至少除以 255。6.2 现象训练集 Dice 很高验证集一塌糊涂原因数据增强只做了训练集验证集分布和训练集差太多或者训练集和验证集有重复图片模型在背答案。解决验证集只做 resize 和归一化不做随机变换。划分数据集时按患者或场景划分不要随机切同一患者的图不能同时出现在训练和验证里。6.3 现象R2UNet 训练几个 epoch 后 loss 变 NaN原因循环结构梯度累积学习率偏高时容易爆。或者某张图掩码全黑Dice 分母为零。解决加梯度裁剪clip_grad_norm_(max_norm1.0)学习率降到 5e-4 或更低。Dice 计算加eps1e-6防止除零。6.4 现象Attention-UNet 报维度不匹配错误原因注意力门控里门控信号和跳跃特征尺寸差一两个像素通常是上采样后尺寸没对齐。解决在AttentionGate里加interpolate对齐或者确保输入图像尺寸是 16 的倍数这样每层下采样和上采样都能整除。6.5 现象显存不够batch_size 降到 1 还是 OOM原因R2UNet 和 Attention-UNet 的计算图比 UNet 大尤其是注意力门控的中间特征。解决先降base到 16 或 32再降输入分辨率到 128。用torch.cuda.empty_cache()清理缓存。混合精度训练torch.cuda.amp能省一半显存但要注意 Dice 计算时转回 float32。7. 从跑通到落地用 ONNX 导出和推理验证收尾三个模型跑通只是起点真正落地要过导出和推理这一关。PyTorch 转 ONNX 是常见做法方便部署到不同推理引擎。导出时最容易翻车的是动态轴没设对导致换输入尺寸就报错。import torch model AttentionUNet(in_ch3, num_classes2).eval() dummy torch.randn(1, 3, 256, 256) torch.onnx.export( model, dummy, attention_unet.onnx, input_names[input], output_names[output], # 把 batch 和宽高设为动态轴换尺寸不用重新导出 dynamic_axes{input: {0: batch, 2: h, 3: w}, output: {0: batch, 2: h, 3: w}}, opset_version12, )逻辑说明dynamic_axes把 batch 和空间维度标为动态推理时可以喂不同尺寸的图。opset_version12兼容性较好太低不支持某些算子太高部分推理引擎不认。导出后用onnxruntime跑一遍对比 PyTorch 输出和 ONNX 输出的最大差值超过 1e-3 说明有算子对不齐。import onnxruntime as ort import numpy as np sess ort.InferenceSession(attention_unet.onnx) inp np.random.randn(1, 3, 256, 256).astype(np.float32) onnx_out sess.run(None, {input: inp})[0] with torch.no_grad(): torch_out model(torch.from_numpy(inp)).numpy() print(最大差值:, np.abs(onnx_out - torch_out).max())参数说明sess.run的第一个参数是输出名列表None表示取全部输出。对比时注意 PyTorch 输出是 logitsONNX 也是不要一边 softmax 一边不 softmax。验证方法上我习惯在验证集里挑几张 Dice 最低的图单独看把原图、标签、预测叠在一起可视化。如果预测边界比标签胖一圈说明损失函数对边界惩罚不够可以加 Dice Loss 和交叉熵联合训练。如果预测漏掉小目标检查下采样次数是不是太多小目标在深层特征里已经没了。我自己的习惯是每换一个数据集先用 UNet 跑一个基线记下 Dice再换 R2UNet 和 Attention-UNet 对比。不要一上来就冲最复杂的基线跑不通复杂模型只会让你更迷茫。这套代码我在广告牌分割和医学影像上都跑过UNet 稳R2UNet 适合深网络Attention-UNet 适合目标稀疏的场景选哪个取决于你的数据长什么样。希望帮到你。本文还有配套的精品资源点击获取
返回列表