简介图像分割中常用的UNet、注意力UNet、残差UNet及两者结合的变体以可运行工程形式打包附带ISIC 2017皮肤病变数据集子集。面向深度学习初学者和医疗影像分析研究者省去自行搭建模型与寻找数据的麻烦方便直接对比四种网络在分割任务上的表现。压缩包共211个文件包括97张PNG和92张JPG皮肤镜图像作为训练与验证样本8个Python脚本覆盖模型定义、训练和预测另有XML配置、Pyc缓存、Shell运行脚本及Markdown说明文档整体约25.48MB目录结构清晰便于复现和改造。目前已有2155人学习下载。资源围绕UNet的编码器解码器结构展示了注意力门控、残差连接以及两者融合的改进方式并针对ISIC数据集配置了预处理、归一化、损失函数与优化器能够帮助读者理解不同机制对分割精度的影响快速开展皮肤病变区域的实验。1. 四款Unet变体加数据集一把梭这套代码到底能帮你省多少事做图像分割的工程师基本都遇到过同一个尴尬论文里四个模型对比写得漂亮真到自己复现时数据格式对不上、训练脚本报错、参数全靠猜一周时间砸进去连个基线都跑不出来。标题里这组组合——Unet、AttentionUnet、R2Unet、R2AUet——正好是分割任务里从入门到进阶的经典路线再配上一份能直接用的数据集想干的就是把「从零复现」变成「从跑通开始」。这套东西适合两类人一类是刚接触分割任务、想弄清楚四个模型到底差在哪的初学者另一类是已经有业务数据、需要快速对比不同骨干做选型的工程师。你不需要自己找数据、写数据加载、调损失函数代码里已经把这些脏活做完了。你只需要改几个参数就能在同一份数据上横向对比四种结构的分割效果这个起点比大多数人想象的省力得多。2. 先把数据对清楚这套分割代码里的数据集结构和预处理逻辑2.1 目录结构与数据流train、mask、val 之间怎么对应拿到代码包先别急着训练第一步是把数据目录结构摸清楚。常见的组织方式是一张原始图对应一张同名 mask 图放在不同子目录下我用 tree 看一眼就明白了dataset/ ├── train/ │ ├── images/ │ │ ├── 0001.png │ │ └── 0002.png │ └── masks/ │ ├── 0001.png │ └── 0002.png └── val/ ├── images/ │ └── 0003.png └── masks/ └── 0003.png这套结构的关键在于文件名严格一一对应。数据加载时最常见的问题是 mask 和 image 名字对不上一旦代码里做了排序拼接名字错位会导致模型拿 A 图画 B 图的标签训练损失函数照样下降但验证集 Dice 永远上不去。我一般会在数据加载器里加一个断言在读取时直接检查文件名是否一致。# data_loader.py import os from torch.utils.data import Dataset from PIL import Image class SegDataset(Dataset): def __init__(self, img_dir, mask_dir, size(256, 256)): self.img_paths sorted([os.path.join(img_dir, f) for f in os.listdir(img_dir)]) self.mask_paths sorted([os.path.join(mask_dir, f) for f in os.listdir(mask_dir)]) # 逐对检查文件名一致防止 train/val 顺序错位 for img_p, mask_p in zip(self.img_paths, self.mask_paths): assert os.path.basename(img_p).split(.)[0] os.path.basename(mask_p).split(.)[0], \ f文件名不匹配: {img_p} vs {mask_p} self.size size这个断言在数据量大的时候会拖慢启动速度但值得。文件名不匹配这个问题我踩过不止一次尤其从网上下载的数据集命名风格五花八门有的带前缀有的不带排序之后很容易错位。另一个需要确认的是 mask 像素值范围数据集来源不同取值也不同有的分割标签是 0 和 1有的是 0 和 255甚至可能是 0 和 65535这个直接决定后面损失函数怎么设计。2.2 预处理脚本灰度图、三通道和 resize 的统一分割任务里最烦的预处理坑是通道数不一致。原始图通常是 RGB 三通道但 mask 是单通道灰度图如果直接统一走 Image.open()PIL 会把单通道的 mask 也读成三通道模型输出和标签对不上。我一般会在预处理阶段把 mask 显式转成单通道同时做归一化# preprocess.py import numpy as np from PIL import Image def load_pair(img_path, mask_path, size(256, 256), mask_threshold127): # 原图保持 RGB img Image.open(img_path).convert(RGB).resize(size, Image.BILINEAR) # mask强制转成单通道灰度再按阈值转成 0/1 mask Image.open(mask_path).convert(L).resize(size, Image.NEAREST) mask_np np.array(mask) mask_bin (mask_np mask_threshold).astype(np.uint8) # 255 - 1 return np.array(img) / 255.0, mask_bin这段代码重点在两个地方。第一resize 的插值方式必须区分原图用双线性mask 用最近邻。如果 mask 也用双线性边缘会产生中间灰度值比如 0.3、0.7 这种训练时会被当成新的类别或者引入噪声。第二mask 的阈值转换把 255 转成 1是为了匹配二分类的标签需求。这套代码如果是做多类分割阈值转换就不适用了得改成映射表方式。还有一点容易被忽略的是 resize 之后 mask 的质量。如果原始标注是在大图上手工画的精细轮廓缩到 256 之后细小的裂缝或者血管可能直接断掉。我处理这类情况一般保持 256 输入但确认一下数据集原始分辨率如果原始图就是 512 甚至 1024 的直接压到 256 会丢掉大量边界细节这四个模型的精度差距也会被缩小因为难点特征都没了。2.3 数据增强选到哪个程度小样本分割的过拟合线很多开源分割数据集的规模不大比如医学场景常用的息肉分割数据集、视网膜血管数据集训练集可能只有几百张。这个量级直接硬训很容易过拟合验证集指标上不去。增强策略我一般用随机翻转加旋转加轻度亮度扰动但有个原则不要做会让目标形态失真的增强。# augmentation.py import random import numpy as np from PIL import Image def aug_pair(img, mask): # 随机水平翻转保证原图和 mask 同步 if random.random() 0.5: img img.transpose(Image.FLIP_LEFT_RIGHT) mask mask.transpose(Image.FLIP_LEFT_RIGHT) # 随机旋转 90 度的倍数保持边缘对齐 k random.choice([0, 1, 2, 3]) img img.rotate(k * 90, resampleImage.BILINEAR) mask mask.rotate(k * 90, resampleImage.NEAREST) # 轻微亮度抖动只作用于原图 if random.random() 0.5: img_np np.array(img).astype(np.float32) img_np * random.uniform(0.9, 1.1) img Image.fromarray(np.clip(img_np, 0, 255).astype(np.uint8)) return img, mask增强代码的核心原则是原图和 mask 必须经历完全一致的几何变换。水平翻转和旋转孤度我都保持同步亮度抖动只作用于原图因为 mask 是标签亮度对它没有意义。旋转角度我限定在 90 度的倍数这样 mask 不用做插值像素对齐关系完全保留。如果用了任意角度的旋转mask 也必须做同样的插值而且插值方式要选最近邻否则边界会出现伪影。增强强度上不要贪随机裁剪也可以考虑但代价是目标可能被切掉一半。对细长型目标比如裂缝、血管这类任务我一般不开随机裁剪翻转加旋转就够了。另外注意增强是每个 epoch 在线做还是离线扩容。在线做的好处是每个 epoch 看到的样本都不同等价于更多训练数据代码里一般放在 Dataset 的getitem里。3. 拆开四个模型Unet、AttentionUnet、R2Unet、R2AUet 的模块差异与适用场景3.1 Unet 基线跳跃连接和特征拼接的起点Unet 的结构图流传很广核心就两个关键词编码器下采样、解码器上采样、跳跃连接。编码器每一层卷积之后做下采样把空间尺寸减半、通道数翻倍逐层提取高层次语义特征解码器反过来逐步恢复空间分辨率跳跃连接把编码器每一层的细节特征直接拼到解码器对应层。# unet.py 核心跳跃连接部分 class Down(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.block 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.block(x) class Up(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.up nn.ConvTranspose2d(in_ch, in_ch // 2, 2, stride2) 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, x1, x2): x1 self.up(x1) # 跳跃连接编码器特征 x2 与上采样特征拼接 diffY x2.size()[2] - x1.size()[2] diffX x2.size()[2] - x1.size()[2] x1 F.pad(x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2]) x torch.cat([x2, x1], dim1) return self.conv(x)这段代码是 Unet 最常见的实现方式。Up 模块接收两个输入x1 是上一层的解码特征x2 是编码器对应层的输出拼接之后通道数正好是 x1 的两倍。如果输入尺寸不是 16 的整数倍上下采样之后尺寸会有 1 到 2 个像素的偏差F.pad 就是用来对齐的。跳跃连接的直观作用是让解码器同时看到高层语义和底层细节底层细节对边缘分割特别关键。这套代码里 Unet 就是基线其他三个模型都是在它的框架上做模块级替换。3.2 AttentionUnet注意力门控挂在哪一层、解决什么问题AttentionUnet 在 Unet 的跳跃连接处加了一个注意力门控Attention Gate。这个门控模块的作用是对编码器传来的特征做加权让模型重点关注和当前目标区域相关的部分抑制背景响应。对医学分割这类目标占比较小的任务这个机制能明显减少误分割。# attention_unet.py 注意力门控 class AttentionGate(nn.Module): def __init__(self, in_ch, g_ch, inter_ch): super().__init__() self.Wg nn.Conv2d(g_ch, inter_ch, 1) self.Wx nn.Conv2d(in_ch, inter_ch, 1) self.psi nn.Conv2d(inter_ch, 1, 1) self.relu nn.ReLU(inplaceTrue) self.sigmoid nn.Sigmoid() def forward(self, x, g): # g 是解码器特征门控信号x 是编码器跳跃连接特征 g1 self.Wg(g) x1 self.Wx(x) # 相加融合后算出注意力权重 out self.relu(g1 x1) out self.sigmoid(self.psi(out)) return x * out注意力门控的输入有两个编码器的跳跃连接特征 x 和解码器当前层的特征 g。g 携带的是更高层的语义信息知道目标大概在哪用它来引导 x 的空域注意力。inter_ch 是中间通道数一般是 min(in_ch, g_ch) 或者直接取 g_ch。训练时可以观察门控输出的热力图如果注意力权重在目标区域之外也有高响应说明门控没学好可以加大中间层通道数或者增加训练轮次。AttentionUnet 对背景复杂、目标区域占比小的分割任务提升明显但对目标本身就很大的场景提升幅度有限因为模型不缺乏定位能力。3.3 R2Unet 和 R2AUet循环残差到底在循环什么R2Unet 的核心是把编码器解码器里的普通卷积块替换成循环残差卷积块Recurrent Residual Convolution Block。普通卷积只做一次卷积操作就输出循环残差块会做多次卷积每次把上一轮的输出重新输入这样同一个卷积层被反复使用多次相当于加深了对同一区域的特征提取。# r2unet.py 循环残差模块 class RecurrentConvBlock(nn.Module): def __init__(self, in_ch, out_ch, t2): super().__init__() self.t t # 循环次数 self.conv nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue) ) def forward(self, x): # 第一次卷积后面每次把上一轮输出重新输入卷积 x1 self.conv(x) for _ in range(self.t - 1): x1 self.conv(x1) # 残差连接输入与循环输出直接相加 return x x1这里有个关键细节循环残差块里的卷积权重是共享的同一个卷积层被反复调用 t 次不是堆叠 t 个不同卷积层。权重共享意味着参数量不增加但计算量增加 t 倍。t 一般取 2取 3 就有明显的显存和算力压力。残差连接解决了循环加深带来的梯度问题让多次卷积不会退化。R2AUet 就是在 R2Unet 的基础上把注意力门控也加回去。R2Unet 负责更充分的特征提取AttentionGate 负责跳跃连接处的特征筛选两者叠加就是 R2AUet 的完整结构。从效果上看R2Unet 单独用会产生大量冗余特征背景区域也被反复强化加上注意力机制后能把这个副作用压住。所以实际使用时如果数据量小、目标占比低我优先选 R2AUet 而不是 R2Unet训练稳定性更好。3.4 四模型对比参数量、训练开销和分割边界表现四个模型的理论边界需要说清楚。Unet 是基线表现稳定但上限最低AttentionUnet 在定位上更好对目标占比小的场景提升明显R2Unet 特征提取上更深但对噪声更敏感训练不稳定R2AUet 综合两者理论上最优但显存占用和训练时间也是最高的。从参数角度AttentionUnet 增加的门控模块参数很少主要在 1x1 卷积整体跟 Unet 接近。R2Unet 循环卷积权重共享单层参数量不变但计算量翻倍。显存方面 R2AUet 因为既有循环残差又有注意力门控占用最高我用 1080Ti 或 2080Ti 这类 11G 显存卡跑 256 输入、batch 8 基本是上限。从分割效果看边界细节是四个模型差距最明显的地方。Unet 的边界偏模糊AttentionUnet 能去掉部分背景噪声R2Unet 对细长目标恢复得更好但对噪声敏感R2AUet 在边界精度和噪声抑制之间平衡得最好。如果你是在做裂纹检测、息肉分割这类细长或小目标任务R2AUet 通常值得第一个试。4. 训练脚本的参数清单从 0 到 1 跑通一次完整实验4.1 数据加载和像素值归一化0/255 与 0/1 的坑提前绕开分割训练跑不通最常见的翻车原因是数据进入模型前的取值不对。很多预训练模型要求输入归一化到 0 到 1 或特定均值方差如果原始图像像素值 0 到 255 直接喂进去模型计算出来的损失和梯度都会异常。下面的加载代码是一个完整的 torch Dataset 实现# dataset.py 完整数据加载 from torch.utils.data import Dataset import torch from PIL import Image import numpy as np class SegDataset(Dataset): def __init__(self, image_dir, mask_dir, img_size256, transformNone, mask_squeezeTrue): self.image_paths sorted(glob.glob(os.path.join(image_dir, *.png))) self.mask_paths sorted(glob.glob(os.path.join(mask_dir, *.png))) self.img_size img_size self.transform transform self.mask_squeeze mask_squeeze def __getitem__(self, idx): img_path self.image_paths[idx] mask_path self.mask_paths[idx] # 原图转 RGBmask 强制单通道 image Image.open(img_path).convert(RGB) mask Image.open(mask_path).convert(L) # 统一尺寸mask 用最近邻 image image.resize((self.img_size, self.img_size), Image.BILINEAR) mask mask.resize((self.img_size, self.img_size), Image.NEAREST) # 转成 tensor 并归一化 image torch.from_numpy(np.array(image)).permute(2, 0, 1).float() / 255.0 mask torch.from_numpy(np.array(mask)).float() / 255.0 if self.mask_squeeze: mask mask.unsqueeze(0) # (1, H, W) return image, mask def __len__(self): return len(self.image_paths)这段代码里我做了两个关键约定图像除以 255 转成 0 到 1 浮点数mask 也除以 255。如果 mask 原始值是 0 和 255除完变成 0 和 1阈值就过了。如果原始值已经是 0 和 1再除以 255 就会把标签变成 0 和 0.004损失函数直接崩。所以拿到数据集先看一眼 mask 的最大值再决定用不用这个除法。不少开源数据集的 mask 是 0 和 255 存储就是为了肉眼查看方便这是最容易被忽略的初始化坑。另外注意 mask_squeeze 这个参数。模型输出是单通道 logitsshape 是 (B, 1, H, W)mask 必须保持同样的 (B, 1, H, W) 才能算损失。如果忘加这个维度PyTorch 广播机制会帮你自动扩展但方向错了损失会算成整个 batch 的混合值指标看着正常实际全错。4.2 Loss 组合和优化器DiceBCE 的权重与常用参数分割任务的损失函数我一般不用单一的交叉熵。医学分割和目标占比小的场景里类别极度不平衡背景像素占比可能超过 95%单纯的 BCE 会让模型学到「全都预测成背景」这种偷懒解。常见做法是 Dice Loss 和 BCE 组合Dice 管区域重合度BCE 管像素级概率。# loss.py Dice BCE 组合损失 import torch import torch.nn as nn class DiceBCELoss(nn.Module): def __init__(self, weight_bce0.5, weight_dice0.5): super().__init__() self.weight_bce weight_bce self.weight_dice weight_dice def forward(self, pred, target): # pred: (B, 1, H, W) 未经过 sigmoid 的 logits # target: (B, 1, H, W) 取值 0/1 bce nn.functional.binary_cross_entropy_with_logits(pred, target) pred torch.sigmoid(pred) smooth 1.0 intersection (pred * target).sum(dim(2, 3)) dice 1 - (2 * intersection smooth) / (pred.sum(dim(2, 3)) target.sum(dim(2, 3)) smooth) dice dice.mean() return self.weight_bce * bce self.weight_dice * dice这里有三个点讲清楚。第一BCE 使用 binary_cross_entropy_with_logits输入是未过 sigmoid 的输出PyTorch 内部做 sigmoid 并计算损失数值稳定性更好。如果先手动 sigmoid 再用普通 BCE在极端概率值下会产生 NaN。第二Dice Loss 的计算里 smooth 取 1.0防止分母为 0。smooth 太小在训练初期容易梯度爆炸。第三两个损失的权重默认各 0.5如果目标占比特别小比如息肉分割我会把 dice 权重提到 0.7。调权重的经验是看验证集 Dice 和准确率的平衡Dice 偏低就加大 dice 权重。优化器上AdamW 是稳妥选择。学习率初始 1e-4配合 ReduceLROnPlateau 在验证指标停滞时下降这是分割训练里最常用的配置比固定学习率省去反复试。# train.py 优化器配置 optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-4) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemax, factor0.5, patience10, verboseTrue )weight_decay 取 1e-4 前后太大容易欠拟合太小起不到正则作用。patience 设 10 表示 10 个 epoch 验证指标不涨就降一半学习率这个是默认稳妥值。如果你的模型训练不稳定可以把 batch size 或学习率同时降下来不要单独只调一个。4.3 评估指标与保存逻辑Dice、IoU 和 best model 的判定分割任务里常用的两个评估指标是 Dice 系数和 IoU。Dice 偏向区域重合度IoU 对边界误差更敏感。很多代码包里这两个都实现了但要注意阈值处理。模型输出是概率图评估前需要先转成 0/1 预测再用 numpy 计算指标这个步骤不能省# eval.py 评估逻辑 import numpy as np def iou_score(pred_mask, true_mask, threshold0.5): # pred_mask: (H, W) 概率值true_mask: (H, W) 0/1 pred_bin (pred_mask threshold).astype(np.uint8) true_bin true_mask.astype(np.uint8) intersection np.logical_and(pred_bin, true_bin).sum() union np.logical_or(pred_bin, true_bin).sum() return intersection / (union 1e-6) def dice_score(pred_mask, true_mask, threshold0.5): pred_bin (pred_mask threshold).astype(np.uint8) true_bin true_mask.astype(np.uint8) intersection np.logical_and(pred_bin, true_bin).sum() return 2 * intersection / (pred_bin.sum() true_bin.sum() 1e-6)评估必须在每个 epoch 结束后做不能只在最后做一次。保存模型时用验证集 Dice 作为标准Dice 最高时保存参数。很多代码给出的是保存最后一个 epoch这在小数据集上有风险最后一轮不一定是最优点。# train.py 模型保存逻辑 best_dice 0.0 for epoch in range(epochs): # training loop... val_dice evaluate(model, val_loader) if val_dice best_dice: best_dice val_dice torch.save(model.state_dict(), fcheckpoints/{model_name}_best.pth) # 顺便保存最新权重方便中断后恢复 torch.save(model.state_dict(), fcheckpoints/{model_name}_last.pth)同时保存 best 和 last 两份权重是必要的。best 是评估最优的模型last 是训练结束时的状态。有时候 last 可能在最后几个 epoch 已经过拟合导致指标下降所以回归对比实验时最好用 best。看哪家模型指标高应该都用各自 best 权重评估这样才是公平对比。4.4 训练参数速查表批次、学习率、epoch、图像尺寸怎么定给一组可以直接上手的参数这不是唯一解但按这套起步基本不会有方向性错误。参数推荐值说明图像尺寸256x256兼顾分辨率和显存原图分辨率高可试 512Batch Size811G 显存不够则降到 4不要低于 2初始学习率1e-4AdamW 配合太高容易震荡学习率调整ReduceLROnPlateau验证指标停滞 10 轮下降一半Epoch100小数据集按 early stop 判断Early Stoppatience 20连续 20 轮验证指标不涨则停止权重初始化HeReLU不要用默认随机初始化batch size 和学习率需要联动。如果调大 batch size 到 16学习率也应该按比例略微提高否则收敛变慢。图像尺寸从 256 改成 512显存占用近似变成四倍这个时候 batch 必须减半减半再减半。epoch 设 100 对大多数小数据集足够剩下的交给 early stop。训练过程里如果 loss 一直震荡降不下去优先排查学习率是不是高了而不是急着换模型结构。5. 运行避坑我踩过的 4 个分割训练常见报错和翻车现场5.1 图像尺寸不是 16 的倍数下采样之后特征图对不齐现象训练时报错错误信息类似size mismatch for x1: got other size。原因Unet 类模型编码器通常下采样 4 次每层尺寸减半两次所以输入尺寸必须是 16 的整数倍。如果原图是 300x300下采样到 18x18再上采样到 288x288跳跃连接拼特征图时尺寸就对不上。虽然我前面的代码里用 F.pad 做了对齐有些实现是直接 assert 尺寸一致训练直接崩。解决所有图像统一 resize 到 16 的整数倍再进模型。数据集里有尺寸不一的图写一个预处理脚本先统一转换然后检查转换后的尺寸列表是否全部符合要求。5.2 标签 mask 被当成了三通道彩色图现象训练正常启动但损失函数永远忽大忽小验证 Dice 一直接近 0。原因数据加载器用Image.open(mask).convert(RGB)读了 mask导致标签变成三通道的彩色数据每个通道值可能相同也可能不同跟模型输出的单通道 logits 对不上。PyTorch 广播会把两个张量强行对齐损失算出来但没有意义。解决mask 一律convert(L)强转单通道灰度再用阈值转成 0/1。拿到新数据集我第一步就是打印 mask 的 shape 和最大值这个习惯能省很多排查时间。5.3 验证集和训练集数据泄漏评估分数虚高现象验证集 Dice 达到 0.95 以上但模型实际分割效果肉眼看着很一般。原因数据集划分时有重叠。常见情况是数据增强只加了训练集但验证集是从增强池里切出来的某些验证样本跟训练样本几乎一样或者随机划分时忘记固定随机种子同一张图同时出现在 train 和 val 目录。解决划分数据前固定随机种子划分完后检查文件名重叠情况# bash 检查 train/val 是否有同名文件 comm -12 \ (ls dataset/train/images | sort) \ (ls dataset/val/images | sort)输出为空才是正常。有输出就说明泄漏了重新划分。分割任务里指标虚高比指标低更危险因为它会让你误判模型能力上线后翻车更狠。5.4 Loss 下降到某个值不动分割结果全是同一块背景现象训练到中后期 Loss 停在某个值附近验证集 Dice 很低预测图全是一个背景色。原因这是类别不平衡的经典表现。模型发现全预测成背景也能拿到很低的损失尤其是在 BCE 权重过高时。Dice Loss 的梯度在这种局面下不够强模型陷入了局部最优。解决提高 Dice 权重到 0.7 以上或者换用 Focal Loss。还有一个土办法把背景像素做下采样从数据层面缓解不平衡。代码里损失函数的 smooth 参数也可以适当调小比如从 1.0 降到 0.5梯度会更敏感。5.5 显存不够patch 训练和 batch 取舍现象1080Ti 上 batch 8、输入 256 直接 CUDA Out Of MemoryR2AUet 尤其严重。原因循环残差模块的计算会多次经过同一卷积层中间特征图占用的显存随循环次数翻倍。AttentionGate 虽然是 1x1 卷积也会增加显存占用。如果模型、数据、batch 三者同时拉满爆显存很正常。解决先降 batch 到 4不行再降输入到 192 或 128。R2Unet 的循环次数 t 从默认 2 降到 1或者换成梯度累积来模拟更大的 batch。# train.py 梯度累积的伪代码写法 accumulation_steps 2 # 等效 batch 翻倍 loss loss / accumulation_steps # 先除以步数 loss.backward() if (step 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()这个技巧比硬调 batch 灵活。梯度累积本质上是把几次反向传播的梯度累加后再做一次参数更新效果接近大 batch。注意 loss 要除以累积步数否则梯度值放大学习率需要相应降低。6. 四个模型都跑通之后用同一个 checkpoint 规则逼出最优结构四个模型跑完一轮之后你手里会有一堆权重文件和评估记录。这时候要做的不是比谁最高就直接用而是把评估口径统一。我常用的做法是每个模型用自己 best 权重做一次全量验证集推理把 Dice、IoU、准确率三个指标拉成一张表然后重点看 IoU 而不是 Dice。Dice 在目标占比小时虚高明显IoU 更保守两者差距越大说明边界噪声越明显。分割效果的可视化验证也不能省。我习惯把预测图叠加在原图上保存成一张图左边原图、中间 mask、右边预测这样一台看下来边界细致程度一目了然。你选模型的时候应该有这样一个判断顺序先看 IoU 够不够再看边界是否连续最后看训练成本能不能接受。R2AUet 如果 IoU 只比 Unet 高 1 到 2 个点但训练时间翻倍那我大概率还是用 Unet 上生产换个更好的数据增强更划算。如果要进一步压榨精度优先试的是 TTATest Time Augmentation。推理时把输入翻转、旋转几次把多个预测概率平均后再做阈值分割。这个方法不用改模型对分割边界有明显改善尤其是不规则形状的目标。代码量也很小十几行就能实现。# tta.py 推理时增强 def predict_tta(model, image, flipsTrue, rotations[0, 90, 180, 270]): model.eval() probs [] with torch.no_grad(): for angle in rotations: img_aug torch.rot90(image, kangle // 90, dims[2, 3]) if flips: img_aug_flip torch.flip(img_aug, dims[3]) for img_in in [img_aug, img_aug_flip]: out torch.sigmoid(model(img_in.unsqueeze(0))) # 反变换回原方向 out torch.flip(out, dims[3]) out torch.rot90(out, k-angle // 90, dims[2, 3]) probs.append(out) else: out torch.sigmoid(model(img_aug.unsqueeze(0))) out torch.rot90(out, k-angle // 90, dims[2, 3]) probs.append(out) return torch.stack(probs).mean(dim0)TTA 对 R2AUet 这种复杂模型涨点效果比 Unet 更明显因为 R2AUet 的预测概率图本身更多样化平均后噪声更少。代价是推理时间成倍增加如果模型要部署到线上TTA 是否值得就看业务要求了。最后一个忠告是保存实验记录。我早期跑对比实验时常犯的错是跑完觉得模型不行就删了权重后来发现是当时学习率没调好后悔药都没有。建议每个模型的每个实验版本都记下参数和指标哪怕看起来是失败实验。这套代码四个结构本身就能玩出很多排列组合数据增强、损失权重、TTA 开关每变一个因素就是一个新实验。记录做得细后面写报告或者调优时省回来的时间远超记录那几分钟。希望帮到你。本文还有配套的精品资源点击获取
阅读完成 · 觉得有帮助?