PyTorch实现U-Net系列:R2U-Net与Attention U-Net实战代码详解
2026/9/16 5:05:32 网站建设 项目流程

简介:这是一种基于PyTorch框架实现的语义分割模型合集,覆盖U-Net、R2U-Net、Attention U-Net、Attention R2U-Net四种主流网络。资源面向图像分割方向的学生、科研人员以及需要快速构建分割基线的开发工程师,能够帮助大家从代码层面理解不同网络结构的差异与改进思路。压缩包内共8个文件,包含7个Python脚本和1个Markdown说明文档。Python脚本按功能拆分为模型定义、数据读取、训练求解、验证评估、工具函数等模块,模块划分清晰,可以直接复用或针对自己的任务进行修改;说明文档则介绍了环境配置、训练步骤与注意事项,对新手比较友好。整个压缩包只有12KB,属于轻量级代码资料,下载和查看都很方便。目前已有555人学习或下载,热度不错。借助这套源码,可以快速对比U-Net与加入循环残差、注意力机制后的各变体在实际效果上的差别,深入理解注意力门控、循环卷积等设计对分割性能的影响,也能为后续网络改造和论文实验提供参考。

1. 医学分割里最稳的基线组合,为什么值得花时间复现一遍

图像分割任务里,U-Net 几乎是所有人绕不开的第一个完整模型。它的编码器-解码器结构配合跳跃连接,在标注样本很少的情况下依然能收敛到可用的分割结果,这也是它从医学影像一路火到遥感、工业质检的原因。项目标题里的 R2U-Net、Attention U-Net、Attention R2U-Net,则是 U-Net 在循环残差和注意力方向上的三种经典改型,而不是互相独立的模型。把它们放在同一份源码里对比训练,能一次性看清结构改动对指标和收敛速度的影响,这对做医疗影像、桥墩病害检测、施工安全场景分割的工程师来说,是最省时间的选型路径。

适合看这篇文章的人有两类:一类是刚接触 PyTorch 基础框架、想用 u-net 复现跑通第一个分割模型的新手;另一类是已经跑过不少分割实验、想搞清楚 Attention Gate 和循环残差到底缓解了什么问题的熟手。下面从网络结构拆起,落到可运行的 PyTorch 代码、训练参数和数据集组织方式,最后给出一套验证和可视化技巧。

2. 从 U-Net 到 R2U-Net:循环残差怎么提升分割精度

这一章先立住 U-Net 这条基线,再把 R2U-Net 的核心改动拆开,最后给出可直接替换的 PyTorch 模块代码。

2.1 编码器-解码器与跳跃连接的核心分工

U-Net 的原理可以用一句话概括:编码器逐层下采样,让网络看到越来越大的感受野;解码器逐层上采样,把低分辨率的高层语义恢复成原图尺寸;跳跃连接把编码器同尺度的特征直接拼到解码器,补回下采样丢失的边界细节。只看骨干结构,它和 FCN 差别不大,但跳跃连接让 U-Net 在少量数据下也能保住小目标。

通道数的设计是一个关键参数,常见做法是首层 64 通道,每过一层翻倍到 512 或 1024。这个“64-128-256-512”的配置几乎不需要改动,除非输入图特别小或类别特别多。深层特征负责回答“这是什么”,浅层特征负责回答“在哪”,因此浅层通道少一点不会伤害分割精度,反而能省显存。

2.2 R2U-Net 的循环卷积与残差连接设计

R2U-Net 的改进不在编码器-解码器框架上,而是把普通卷积块换成了 Recurrent Residual CNN 块。循环卷积的意思是在同一个卷积层上反复作用 t 次,让每个位置累积更大的等效感受野;残差连接则让梯度能直接穿过时间步,避免深层网络训练时的梯度消失。原文推荐 t=2,即每个块内做两次循环,这个值不需要盲目加大,t=3 以上收益很小且显存上涨明显。

2.2.1 用 PyTorch 实现循环卷积模块
import torch import torch.nn as nn class RecurrentConv(nn.Module): def __init__(self, in_ch, out_ch, t=2): super().__init__() self.conv = nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1) self.bn = nn.BatchNorm2d(out_ch) self.relu = nn.ReLU(inplace=True) self.t = t def forward(self, x): for _ in range(self.t): x = self.relu(self.bn(self.conv(x))) return x class RRCNNBlock(nn.Module): def __init__(self, in_ch, out_ch, t=2): super().__init__() self.conv1 = RecurrentConv(in_ch, out_ch, t) self.conv2 = RecurrentConv(out_ch, out_ch, t) def forward(self, x): x1 = self.conv1(x) x2 = self.conv2(x1) return x1 + x2 # 残差路径:第一轮输出直接加到第二轮结果

逻辑说明:RecurrentConv里的for循环是时间维度的循环,同一组卷积权重被反复作用,而不是串联多个不同卷积层,这能控制参数量不随 t 线性增大。RRCNNBlockx1 + x2在实现时需要注意维度,in_chout_ch不一致时不能直接相加。常见做法是先对输入做 1x1 卷积对齐通道,再把对齐结果加进去,否则会报 shape mismatch。

2.3 编码器与解码器的整体替换策略

R2U-Net 替换 U-Net 时不需要改动跳跃连接和图上下采样逻辑,只需把每层的 double conv 换成RRCNNBlock。下面这份代码是编码器和解码器的最小骨架,可以直接跑在 256x256 的单通道输入上。

class EncoderBlock(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.rr = RRCNNBlock(in_ch, out_ch) self.pool = nn.MaxPool2d(2) def forward(self, x): return self.rr(x), self.pool(self.rr(x)) class DecoderBlock(nn.Module): def __init__(self, in_ch, skip_ch, out_ch): super().__init__() self.up = nn.ConvTranspose2d(in_ch, out_ch, kernel_size=2, stride=2) self.rr = RRCNNBlock(out_ch + skip_ch, out_ch) def forward(self, x, skip): x = self.up(x) x = torch.cat([x, skip], dim=1) return self.rr(x)

参数说明:EncoderBlock.forward里调用了两次self.rr(x),第一次输出用于跳连,第二次输出继续下采样,这会让显存翻倍,实际工程里更推荐先算out = self.rr(x),再返回outself.pool(out)ConvTranspose2d的 stride 和 kernel_size 必须都为 2,否则输出尺寸和skip对不上,这一步是初学者最容易踩的坑。

模块输入通道输出通道特征图尺寸
Encoder Block 1164256x256
Encoder Block 264128128x128
Encoder Block 312825664x64
Encoder Block 425651232x32
Bottleneck512102416x16

3. Attention U-Net 的门控注意力:从代码看跳连改进

Attention U-Net 的最大贡献是在跳跃连接上加了 Attention Gate,让解码器在融合特征时自动抑制无关区域。这一章会解释门控机制为什么有效,并给出可以直接嵌入模型的实现代码。

3.1 普通跳连在语义不对齐时的困境

编码器浅层特征分辨率高但语义弱,解码器深层特征语义强但分辨率低。普通跳连直接把两者在通道维拼接,等于默认所有位置都同等重要,这在目标区域小、背景杂乱的医学图像上会拉低收敛速度。Attention Gate 的核心在于让解码器生成一个空间权重图,乘到跳连特征上,突出与目标相关的区域,抑制背景响应。

这个权重图不是额外学的分类头,而是由解码器当前层的门控信号 g 和编码器跳连特征 x 共同计算得到。g 包含高层语义,x 包含细节纹理,两者融合后才能正确判断哪些浅层位置值得传递。

3.2 用 PyTorch 实现二维 Attention Gate

class AttentionGate(nn.Module): def __init__(self, in_ch_g, in_ch_x, out_ch): super().__init__() self.W_g = nn.Sequential( nn.Conv2d(in_ch_g, out_ch, kernel_size=1), nn.BatchNorm2d(out_ch) ) self.W_x = nn.Sequential( nn.Conv2d(in_ch_x, out_ch, kernel_size=1), nn.BatchNorm2d(out_ch) ) self.psi = nn.Sequential( nn.Conv2d(out_ch, 1, kernel_size=1), nn.BatchNorm2d(1), nn.Sigmoid() ) def forward(self, g, x): g1 = self.W_g(g) x1 = self.W_x(x) attn = torch.relu(g1 + x1) attn = self.psi(attn) attn = F.interpolate(attn, size=x.shape[-2:], mode='bilinear', align_corners=False) return x * attn

逻辑说明:W_gW_x都用 1x1 卷积把输入压缩到同一个中间通道数out_ch,之后相加并过 ReLU,再用psi把每个空间位置压缩到 0 到 1 之间。interpolate是必要的一步,因为g的分辨率通常小于x,注意力图的尺寸首先要对齐x,才能做逐像素相乘。out_ch在原文中取in_ch_x // 2,如果显存不紧张,可以直接用in_ch_x,效果差异很小。

3.2.1 在解码器中插入 Attention Gate
class AttnDecoderBlock(nn.Module): def __init__(self, in_ch, skip_ch, out_ch): super().__init__() self.up = nn.ConvTranspose2d(in_ch, out_ch, 2, stride=2) self.attn = AttentionGate(in_ch_g=out_ch, in_ch_x=skip_ch, out_ch=skip_ch // 2) self.conv = RRCNNBlock(out_ch + skip_ch, out_ch) def forward(self, x, skip): x = self.up(x) # x 此时是门控信号 g skip = self.attn(g=x, x=skip) # 用 g 加权 skip x = torch.cat([x, skip], dim=1) return self.conv(x)

参数说明:这里in_ch_g用的是上采样后的out_ch,而不是解码器输入的in_ch,因为门控信号需要在空间分辨率上和 skip 更接近,注意力计算才稳定。skip_ch // 2若在极端情况下得到 0,需要改成max(1, skip_ch // 2)

3.3 加 Attention 之后的参数量和显存变化

Attention Gate 每个解码器层增加约 3 个 1x1 卷积,整套模型额外参数量不到 10%,但对显存的影响不能只看参数。注意力图要和 skip 同尺寸存储,所以显存占用会有小幅上涨,大约 5% 到 8%。如果训练时 OOM,优先检查 4 层 Attention Gate 是否可以只在最后两层使用,这不是论文里的标准配置,但实际工程里能保住 batch size,且精度损失常在 0.5% 以内。

提示:在 torch 1.10 之后的版本里,F.interpolatealign_corners参数保持False,否则梯度回传到g1时可能出现棋盘格伪影。

4. Attention R2U-Net 训练配置:模型组合与关键参数

前三章分别实现了 U-Net、R2U-Net 和 Attention U-Net,Attention R2U-Net 就是把 RRCNNBlock 和 AttentionGate 放进同一个解码器里。这一章先讲模型怎么组合,再给出训练时最值得调的 8 个参数,并给出损失函数的实现。

4.1 组合模型的模块复用与整体结构

组合方式非常直接:编码器使用 RRCNNBlock,解码器使用带 Attention Gate 的 RRCNNBlock 解码块。这样循环残差负责加强特征提取的深度,注意力负责优化跳连融合,两者解决的问题不重叠。以下代码展示完整的 Attention R2U-Net 头部与解码器装配流程。

class AttentionR2UNet(nn.Module): def __init__(self, in_ch=1, out_ch=1, t=2): super().__init__() self.enc1 = EncoderBlock(in_ch, 64) self.enc2 = EncoderBlock(64, 128) self.enc3 = EncoderBlock(128, 256) self.enc4 = EncoderBlock(256, 512) self.bottleneck = RRCNNBlock(512, 1024, t) self.dec4 = AttnDecoderBlock(1024, 512, 512) self.dec3 = AttnDecoderBlock(512, 256, 256) self.dec2 = AttnDecoderBlock(256, 128, 128) self.dec1 = AttnDecoderBlock(128, 64, 64) self.out_conv = nn.Conv2d(64, out_ch, 1) def forward(self, x): e1, p1 = self.enc1(x) e2, p2 = self.enc2(p1) e3, p3 = self.enc3(p2) e4, p4 = self.enc4(p3) b = self.bottleneck(p4) d = self.dec4(b, e4) d = self.dec3(d, e3) d = self.dec2(d, e2) d = self.dec1(d, e1) return self.out_conv(d)

结构说明:EncoderBlock.forward返回两个值,第一个是送入跳连的特征,第二个是下采样后的特征,务必保持顺序一致。out_conv使用 1x1 卷积,而不是 3x3,因为最后只需要把通道数压缩到类别数,不需要额外感受野。

4.2 训练阶段的 8 个关键参数设置

这四个模型在同一数据集上的训练配置可以完全一致,差别只体现在收敛速度和最终精度上。建议优先使用下面的参数表作为初始配置,再根据自己的数据调整。

参数推荐值说明
输入尺寸256x256超过 512 时显存压力陡增
初始学习率1e-4比分类任务小一个量级,分割损失面更抖
学习率调度ReduceLROnPlateau监控验证集 Dice,patience=8
优化器AdamWweight_decay=1e-4
batch size8-16小于 4 时优先调小图尺寸而不是改架构
损失函数BCE + Dice权重各 0.5,样本不均时 Dice 权重可提到 0.7
最大 epoch100配合早停,patience=15
梯度裁剪max_norm=12防止 R2 循环展开时梯度爆炸

4.3 损失函数实现与不平衡样本处理

医学图像分割里背景像素往往远多于前景,单独用 BCE 会让网络偏向输出全零图。Dice Loss 直接优化交并比,对小目标和类别不平衡更友好。和 BCE 组合时,要注意DiceLoss内部对predsigmoid,因此输入必须是 logits,不能在外部提前过sigmoid

class DiceLoss(nn.Module): def __init__(self, smooth=1.0): super().__init__() self.smooth = smooth def forward(self, pred, target): pred = torch.sigmoid(pred).reshape(pred.size(0), -1) target = target.reshape(target.size(0), -1) intersection = (pred * target).sum(dim=1) union = pred.sum(dim=1) + target.sum(dim=1) dice = (2 * intersection + self.smooth) / (union + self.smooth) return 1 - dice.mean() def combined_loss(pred, target, w_bce=0.5, w_dice=0.5): bce = nn.functional.binary_cross_entropy_with_logits(pred, target) dice = DiceLoss()(pred, target) return w_bce * bce + w_dice * dice

参数说明:smooth的作用是防止某些训练轮次里前景完全未被预测导致分母为 0,取值 1.0 是常用默认值。reshape(pred.size(0), -1)会把每个样本单独展平,保证 loss 是 batch 内各样本 Dice 的均值,而不是把所有样本混在一起计算,后者在 batch 内分布不均时会掩盖单个样本的退化。

环境方面,PyTorch 2.x 在torch.compile下训练这类小模型收益不大,反而拉长每次前向的时间。建议直接用 pip 或 conda 安装 pytorch 稳定版即可,不需要刻意追踪最新 nightly 构建。GPU 型号较旧时,注意 CUDA 版本和 torch 的 wheel 匹配,优先到 pytorch 官网查自己的 CUDA 版本再选安装命令。

5. 数据集与预处理:用 PyTorch 打通标注到训练的输入管线

标题里“数据集”的分量不轻,很多 u-net 复现卡在最后一步:模型写完了,却没有合适的数据喂进去。这一章给出一套适用于公共分割数据集和自建标注数据的通用预处理流程。

5.1 分割数据集的目录组织与命名规范

推荐的目录结构如下:

datasets/ ├── images/ │ ├── case001.png │ ├── case002.png ├── masks/ │ ├── case001.png │ ├── case002.png ├── train.txt └── val.txt

train.txtval.txt每行放一个不带扩展名的文件名。使用文本索引而不是扫描文件夹,最大的好处是换数据集时不用改 Dataset 代码,只换索引文件。如果做交叉验证,只需生成多份 txt,训练代码完全不变。

命名规范建议统一成小写加下划线,避免 Windows 和 Linux 文件系统之间的兼容问题。mask 中前景区域建议统一为 255,背景为 0,在数据加载时再做归一化到 0 和 1。不要使用边缘抗锯齿的 PNG,否则阈值化会引入噪点。

5.2 自定义 Dataset 类的完整实现

from torch.utils.data import Dataset from PIL import Image import torch import os class SegmentationDataset(Dataset): def __init__(self, img_dir, mask_dir, file_list, transform=None): self.img_dir = img_dir self.mask_dir = mask_dir self.file_list = [line.strip() for line in open(file_list)] self.transform = transform def __len__(self): return len(self.file_list) def __getitem__(self, idx): name = self.file_list[idx] img_path = os.path.join(self.img_dir, name + ".png") mask_path = os.path.join(self.mask_dir, name + ".png") img = Image.open(img_path).convert("RGB") mask = Image.open(mask_path).convert("L") if self.transform is not None: img, mask = self.transform(img, mask) mask = (mask > 0).float() return img, mask

这里的transform需要同时处理 image 和 mask,不能直接使用torchvision.transforms.ToTensor后再编写独立的mask_transform,因为随机翻转时两者的空间对应关系必须保持一致。下面给出一个最简单的兼容组合:

from torchvision import transforms class TrainTransform: def __init__(self): self.flip = transforms.RandomHorizontalFlip(p=0.5) def __call__(self, img, mask): if torch.rand(1) > 0.5: img = self.flip(img) mask = self.flip(mask) img = transforms.ToTensor()(img) mask = transforms.ToTensor()(mask) return img, mask

参数说明:torch.rand(1) > 0.5这个随机决策要和外层RandomHorizontalFlip(p=0.5)解耦,否则图片和掩膜各自随机翻转会对不上。实践中更推荐使用albumentations这类库,它原生保证同一个随机种子同时作用在 image 和 mask 上。

5.3 数据增强、尺度统一与样本均衡

模型的输入尺寸决定了感受野和显存占用,最常见的选择是 256x256 或 512x512。训练时先 resize 到一个略大的尺寸再随机 crop 到目标尺寸,等于免费获得随机裁剪增强。验证时直接 resize 到目标尺寸,不做 crop。

医学图像常用的增强组合为:随机旋转 90 度或 180 度、水平翻转、局部弹性形变、亮度对比度扰动。做分割时慎用光照变化过大的增强,因为有些模态下亮度本身就携带语义信息。弹性形变的强度参数 sigma 取 5 到 10,过大容易让解剖结构扭曲到失真。

提示:如果显存不够但想保证感受野不缩小,优先把首层通道数从 64 降到 32,而不是把输入图缩小到 128。缩图带来的精度损失通常大于减通道。

6. 用 TTA 与指标验证 U-Net 系列模型:快速评测与可视化的实用技巧

最后一章聚焦在验证环节。训练完成后,模型的最终报告不能只看训练集 Dice,要在验证集上做标准评测,并配合简单的测试时增强提升落地的稳定性。

计算验证集指标时,推荐同时输出 Dice 和 IoU,因为两者在不同尺度目标上的敏感度不同。R2U-Net 的循环结构让单次前向推理时间比普通 U-Net 多约 20%,在做大图滑窗推理时要把这部分时间预算算进去。

下面的predict_tta函数实现水平翻转 TTA,适合所有四个模型复用:

import torch import torch.nn.functional as F def predict_tta(net, img, device): net.eval() with torch.no_grad(): x = img.unsqueeze(0).to(device) pred1 = torch.sigmoid(net(x)) pred2 = torch.sigmoid(net(torch.flip(x, dims=[-1]))) pred = (pred1 + torch.flip(pred2, dims=[-1])) / 2 return pred.squeeze(0).cpu()

参数说明:dims=[-1]对应最后一维即宽度方向,flipflip必须配对使用,否则对齐后预测结果相反。TTA 的核心收益是消除翻转前后模型在边界上的系统性偏移,它不能替代数据增强,只会让单次推理的稳定性更好。在验证集上开启 TTA 后,Dice 通常能上涨 0.5 到 1.5 个百分点,但当训练数据本身就是水平翻转增强时,收益会显著缩小。

视觉验证时,建议把原图、预测掩膜、真实掩膜三张图叠在一起保存,每张预测图都留下一份,方便回溯是哪一类错误导致指标下降。使用 matplotlib 保存三通道对比图时,注意把预测概率和真实掩膜同时映射到 0-255 区间,否则灰度图会出现全黑或全白的情况。

关于滑窗推理,U-Net 系列模型对固定尺寸输入最友好。预测大图时,推荐使用带 overlap 的滑窗,overlap 取窗口边长的 1/8 到 1/4。窗口拼回原图时,overlap 区域用重叠次数取平均,而不是直接覆盖,这样可以避免拼接缝。这四个模型在推理阶段都能用torch.cuda.amp.autocast启用半精度,Attention Gate 里的interpolate在 fp16 下表现稳定,但 BatchNorm 层在显存允许时建议保留 fp32,否则小 batch 下统计量偏差会放大。将 TTA 和半精度集成进验证脚本后,用同一份验证集跑一次完整评测,保存原图、TTA 预测和真实掩膜的三图对比结果,再根据 slice 的亮度分布决定是否需要调 loss 权重。

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询