简介:图像分割中常用的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_threshold=127): # 原图:保持 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, resample=Image.BILINEAR) mask = mask.rotate(k * 90, resample=Image.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, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True) ) 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, stride=2) self.conv = nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True) ) 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], dim=1) 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(inplace=True) 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, t=2): super().__init__() self.t = t # 循环次数 self.conv = nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True) ) 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_size=256, transform=None, mask_squeeze=True): 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] # 原图转 RGB,mask 强制单通道 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 这个参数。模型输出是单通道 logits,shape 是 (B, 1, H, W),mask 必须保持同样的 (B, 1, H, W) 才能算损失。如果忘加这个维度,PyTorch 广播机制会帮你自动扩展,但方向错了损失会算成整个 batch 的混合值,指标看着正常实际全错。
4.2 Loss 组合和优化器:Dice+BCE 的权重与常用参数
分割任务的损失函数我一般不用单一的交叉熵。医学分割和目标占比小的场景里,类别极度不平衡,背景像素占比可能超过 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_bce=0.5, weight_dice=0.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(), lr=1e-4, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, mode='max', factor=0.5, patience=10, verbose=True )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, threshold=0.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, threshold=0.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(), f'checkpoints/{model_name}_best.pth') # 顺便保存最新权重,方便中断后恢复 torch.save(model.state_dict(), f'checkpoints/{model_name}_last.pth')同时保存 best 和 last 两份权重是必要的。best 是评估最优的模型,last 是训练结束时的状态。有时候 last 可能在最后几个 epoch 已经过拟合导致指标下降,所以回归对比实验时最好用 best。看哪家模型指标高,应该都用各自 best 权重评估,这样才是公平对比。
4.4 训练参数速查表:批次、学习率、epoch、图像尺寸怎么定
给一组可以直接上手的参数,这不是唯一解,但按这套起步基本不会有方向性错误。
| 参数 | 推荐值 | 说明 |
|---|---|---|
| 图像尺寸 | 256x256 | 兼顾分辨率和显存,原图分辨率高可试 512 |
| Batch Size | 8(11G 显存) | 不够则降到 4,不要低于 2 |
| 初始学习率 | 1e-4 | AdamW 配合,太高容易震荡 |
| 学习率调整 | ReduceLROnPlateau | 验证指标停滞 10 轮下降一半 |
| Epoch | 100 | 小数据集按 early stop 判断 |
| Early Stop | patience 20 | 连续 20 轮验证指标不涨则停止 |
| 权重初始化 | He(ReLU) | 不要用默认随机初始化 |
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 Memory,R2AUet 尤其严重。
原因:循环残差模块的计算会多次经过同一卷积层,中间特征图占用的显存随循环次数翻倍。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 上生产,换个更好的数据增强更划算。
如果要进一步压榨精度,优先试的是 TTA(Test Time Augmentation)。推理时把输入翻转、旋转几次,把多个预测概率平均后再做阈值分割。这个方法不用改模型,对分割边界有明显改善,尤其是不规则形状的目标。代码量也很小,十几行就能实现。
# tta.py 推理时增强 def predict_tta(model, image, flips=True, rotations=[0, 90, 180, 270]): model.eval() probs = [] with torch.no_grad(): for angle in rotations: img_aug = torch.rot90(image, k=angle // 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(dim=0)TTA 对 R2AUet 这种复杂模型涨点效果比 Unet 更明显,因为 R2AUet 的预测概率图本身更多样化,平均后噪声更少。代价是推理时间成倍增加,如果模型要部署到线上,TTA 是否值得就看业务要求了。
最后一个忠告是保存实验记录。我早期跑对比实验时常犯的错是:跑完觉得模型不行就删了权重,后来发现是当时学习率没调好,后悔药都没有。建议每个模型的每个实验版本都记下参数和指标,哪怕看起来是失败实验。这套代码四个结构本身就能玩出很多排列组合,数据增强、损失权重、TTA 开关,每变一个因素就是一个新实验。记录做得细,后面写报告或者调优时省回来的时间远超记录那几分钟。希望帮到你。
本文还有配套的精品资源,点击获取