简介:面向医学图像分割学习者的UNet+DRIVE完整工程包,帮助读者掌握视网膜血管分割的模型搭建、训练与预测流程。包内包含98个文件,以82张png图像为主,涵盖DRIVE数据集与预测结果可视化,另有4个Python源码、4个编译后的pyc文件、1个预训练权重UNet.pth及相关配置与XML文件,压缩包共115.32MB,可支持从数据加载、模型训练到推理预测的完整实践。资源已吸引3668人学习,适合深度学习入门者及医学影像方向研究者参考。通过源码可学习UNet收缩路径、扩展路径与跳跃连接的实现细节,结合预训练模型和结果图可直接观察分割效果;配套讲解还对数据预处理、Dice评估、损失函数优化等给出思路,是动手复现论文模型的实用资料。
1. UNet 网络做图像分割 DRIVE 数据集:先把这件事拆明白
如果你接过眼底图像血管分割的需求,就会知道难点不在模型结构,而在数据。DRIVE 数据集只有 40 张眼底照片,20 张训练、20 张测试,每张都有专家手工标注的血管 mask。用 UNet 网络做图像分割 DRIVE 数据集,是很多人在医学图像分割上迈出的第一步,因为它把“小数据 + 分割网络 + 标准评估”这条链路完整跑通了。我之前在某公司的模拟项目X里就是用这个组合验证分割管线,后来发现真正决定结果的是预处理、损失函数和训练参数,而不是网络本身。这篇文章面向准备在自己机器上复现这个任务的开发者,或者想用 UNet 迁移到其他医学分割场景的人。读完你至少能跑通一个 Dice 0.78 左右的血管分割模型,并知道踩坑时先查哪里。
2. DRIVE 数据集与预处理:40 张图怎么喂给 UNet
2.1 DRIVE 的目录结构与标注含义
DRIVE 数据集的结构非常规整,下载解压后你会看到两个顶层目录:training 和 test。训练集 20 张,测试集 20 张,每张图像都有对应的专家手工标注。以训练集为例,一个样本由三部分组成:原始眼底图、1st_manual 目录下的血管标注、mask 目录下的 FOV(Field of View)区域掩码。FOV mask 的像素值为 255 表示眼底区域,0 表示黑色背景。很多人第一次跑的时候直接拿整张图算损失,结果把大量黑色背景也当成预测目标,导致 Dice 虚高但实际分割很差。
tree DRIVE -L 2输出会是这样的大致结构:
DRIVE/ ├── training/ │ ├── images/ │ ├── 1st_manual/ │ └── mask/ └── test/ ├── images/ ├── 1st_manual/ ├── 2nd_manual/ └── mask/注意测试集比训练集多了一个 2nd_manual 目录,这是第二位专家的标注。DRIVE 的惯例做法是拿 1st_manual 作为标准答案来算评价指标,2nd_manual 可以用于比较专家之间的标注差异。我在模拟项目X里也试过用两位专家的交集做训练目标,但提升并不明显,所以建议你直接使用 1st_manual 当后续训练的 ground truth。
读取时要注意 DRIVE 的图像是 RGB 三通道眼底图,但血管标注和 FOV mask 是单通道灰度图。它们的分辨率一致,都是 565×584。这个尺寸并不是 16 的倍数,而 UNet 通常下采样 4 次,特征图尺寸按 2 的指数倍缩小,所以输入尺寸最好是 16 的倍数,否则你会在拼接跳跃连接时遇到尺寸对不上的隐性报错。
2.2 读图、归一化与数据增强
DRIVE 原图周围有一圈黑色背景,FOV mask 的作用就是把真正有用的眼底区域框出来。我一般会先把图像和 mask 都读成数组,然后统一 resize 到 576×576,这样既满足 16 的倍数要求,又不会像缩到 256 那样丢失太多细血管细节。DRIVE 血管最细的地方只有 1~2 个像素宽,过度下采样会让毛细血管直接消失,后期怎么调损失函数都救不回来。
import cv2 import numpy as np def load_drive_sample(image_path, label_path, fov_path, size=(576, 576)): image = cv2.imread(image_path, cv2.IMREAD_COLOR) image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) label = cv2.imread(label_path, cv2.IMREAD_GRAYSCALE) fov = cv2.imread(fov_path, cv2.IMREAD_GRAYSCALE) image = cv2.resize(image, size, interpolation=cv2.INTER_LINEAR) label = cv2.resize(label, size, interpolation=cv2.INTER_NEAREST) fov = cv2.resize(fov, size, interpolation=cv2.INTER_NEAREST) label = (label > 127).astype(np.float32) fov = (fov > 127).astype(np.float32) image = image.astype(np.float32) / 255.0 return image, label, fov这里有几个关键点。resize 标注时一定要用 INTER_NEAREST,因为血管是二值结构,线性插值会在边界产生灰色过渡层,不仅让 ground truth 变得不准确,还会导致后来的 Dice 计算出现模糊像素。FOV mask 也一样,用最近邻插值可以保持 0/1 的硬边界。图像归一化用简单的除以 255 就够,不需要额外的均值方差标准化,因为后面损失函数对尺度不敏感。
数据增强方面,我推荐用 albumentations 库,它对 mask 同步变换处理得很干净。医学图像分割和自然图像不同,旋转、翻转这类几何增强非常有效,因为血管的走向本身没有固定的上下方向。弹性形变也能模拟眼底血管的个体差异,但不要用过大的强度,否则细血管会被扯断。
import albumentations as A train_transform = A.Compose([ A.RandomRotate90(p=0.5), A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5), A.ElasticTransform(alpha=8, sigma=3, p=0.3), A.RandomBrightnessContrast(brightness_limit=0.1, contrast_limit=0.1, p=0.3), ])ElasticTransform 的 alpha 参数控制形变强度,sigma 控制平滑程度。对 DRIVE 来说,alpha 设 8、sigma 设 3 是比较稳妥的起点。如果调大 alpha 到 20,增强样本的血管会出现明显断裂,模型会把断裂也学进去,测试时预测图就容易出现血管不连续的问题。
2.3 训练样本构造:patch 还是整图
这是新手最容易纠结的问题。DRIVE 只有 20 张训练图,直接整图训练意味着一个 epoch 只有 20 个样本,UNet 很快就会过拟合。常见做法是随机裁剪 patch。patch 尺寸选 256×256,步长 128,这样每张图能裁剪出大约十几个 patch,20 张图就能得到几百个训练样本,如果加上增强,一个 epoch 可以过一千多个 patch。
def extract_patches(image, label, fov, patch_size=256, stride=128): h, w = image.shape[:2] patches_img, patches_lbl, patches_fov = [], [], [] for y in range(0, h - patch_size + 1, stride): for x in range(0, w - patch_size + 1, stride): img_p = image[y:y + patch_size, x:x + patch_size] lbl_p = label[y:y + patch_size, x:x + patch_size] fov_p = fov[y:y + patch_size, x:x + patch_size] if fov_p.mean() < 0.1: continue if lbl_p.sum() < 1: continue patches_img.append(img_p) patches_lbl.append(lbl_p) patches_fov.append(fov_p) return np.stack(patches_img), np.stack(patches_lbl), np.stack(patches_fov)过滤条件很关键。fov_p.mean() < 0.1 表示该 patch 落在眼底区域外的部分太多,这种 patch 对训练没有意义。lbl_p.sum() < 1 表示 patch 内完全没有血管像素,如果大量喂这种样本,模型会倾向于输出全背景,导致漏检小血管。采用步长小于 patch 尺寸的重叠裁剪,能让相邻 patch 之间共享边界信息,训练时模型不会因为血管恰好在 patch 边缘而学不到。
patch 大小会直接影响显存占用。在 1080Ti 这类 11G 显存的卡上,256×256 patch 的 batch size 可以开到 16。如果你只有 6G 显存,建议把 patch 降到 192,或者把 batch size 降到 8。patch 太小会丢失上下文,之前测试过 128×128,Dice 大概低了 2 个百分点,所以不要为了省显存无限缩 patch。
3. UNet 网络结构与关键参数:在 PyTorch 里从零搭一个分割模型
3.1 编码器-解码器与跳跃连接:UNet 为什么适合 DRIVE
UNet 的核心是“编码器收缩、解码器扩张”的对称结构。编码器通过卷积和池化逐步降低特征图分辨率,提取越来越抽象的语义信息;解码器通过上采样把低分辨率的特征图恢复到原图尺寸。如果只是这样,解码器恢复出的边缘会非常粗糙。UNet 的巧妙之处是加入了跳跃连接,把编码器中同尺度的特征图直接拼接到解码器上,相当于给解码器递了一张保留血管边缘细节的小抄。
DRIVE 任务的特殊性在于目标是很细的血管,而且血管与背景的对比度不高。没有跳跃连接,网络在池化过程中会丢失大量血管末梢信息,预测结果往往是主干血管清晰、细血管断裂。跳跃连接让解码器在每个尺度上都能同时看到高层语义和低层边缘,这也是 UNet 在 DRIVE 这类小目标分割上比普通 FCN 稳定的原因。
UNet 原始结构用 3×3 卷积加 ReLU,下采样用 2×2 MaxPooling,上采样用转置卷积。实际复现时不需要完全照搬,可以在每个卷积块后面加 BatchNorm,这样深层网络收敛更快,对学习率容忍度也更高。唯一的代价是 BN 在小 batch size 下统计量会抖动,所以如果你只能用 batch size 4 以下训练,可以把 BN 换成 InstanceNorm。
3.2 核心代码:从零写一个适用于二分类的 UNet
下面这段代码是 DRIVE 训练最常用的一种 UNet 实现,输入 3 通道 RGB,输出 1 通道分割 logits,没有内置 sigmoid,方便用 BCEWithLogitsLoss 组合 Dice Loss。base 参数设置为 64,表示第一层卷积的输出通道数。
import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() 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, x): return self.conv(x) class UNet(nn.Module): def __init__(self, in_channels=3, num_classes=1, base=64): super().__init__() self.enc1 = DoubleConv(in_channels, base) self.pool1 = nn.MaxPool2d(2) self.enc2 = DoubleConv(base, base * 2) self.pool2 = nn.MaxPool2d(2) self.enc3 = DoubleConv(base * 2, base * 4) self.pool3 = nn.MaxPool2d(2) self.enc4 = DoubleConv(base * 4, base * 8) self.pool4 = nn.MaxPool2d(2) self.bridge = DoubleConv(base * 8, base * 16) self.up4 = nn.ConvTranspose2d(base * 16, base * 8, 2, stride=2) self.dec4 = DoubleConv(base * 16, base * 8) self.up3 = nn.ConvTranspose2d(base * 8, base * 4, 2, stride=2) self.dec3 = DoubleConv(base * 8, base * 4) self.up2 = nn.ConvTranspose2d(base * 4, base * 2, 2, stride=2) self.dec2 = DoubleConv(base * 4, base * 2) self.up1 = nn.ConvTranspose2d(base * 2, base, 2, stride=2) self.dec1 = DoubleConv(base * 2, base) self.outc = nn.Conv2d(base, num_classes, kernel_size=1) def forward(self, x): e1 = self.enc1(x) e2 = self.enc2(self.pool1(e1)) e3 = self.enc3(self.pool2(e2)) e4 = self.enc4(self.pool3(e3)) b = self.bridge(self.pool4(e4)) d4 = self.up4(b) d4 = torch.cat([d4, e4], dim=1) d4 = self.dec4(d4) d3 = self.up3(d4) d3 = torch.cat([d3, e3], dim=1) d3 = self.dec3(d3) d2 = self.up2(d3) d2 = torch.cat([d2, e2], dim=1) d2 = self.dec2(d2) d1 = self.up1(d2) d1 = torch.cat([d1, e1], dim=1) d1 = self.dec1(d1) return self.outc(d1)这段代码的关键点在于跳跃连接的 channel 对齐。比如 up4 输出 256 个通道,而 e4 也是 256 个通道,拼接后变成 512 个通道,所以 dec4 的第一个 DoubleConv 输入必须是 512。这也是为什么 down 和 up 路径的通道数设置成倍数关系,严谨地一一对应。如果你修改 base,通道关系会自动等比缩放。
代码里使用 1×1 卷积作为输出层,把 base 通道映射到 num_classes=1。对 DRIVE 来说 num_classes=1 表示只预测血管这一个前景类,背景不需要单独输出一个通道。这种二类分割写法比输出两个通道再用 argmax 更省显存,计算 Dice 也方便。如果你的任务变成多类分割,就把 num_classes 改成类别数,并将损失函数改为 CrossEntropyLoss。
3.3 损失函数与评估指标:Dice、IoU 与 AUC 怎么选
DRIVE 的血管像素占比大约只有 10%~12%,存在严重的类别不平衡。如果直接使用普通交叉熵,模型会发现只要全部预测为背景就能得到很低的 loss,训练出的结果经常是全黑图。所以损失函数要加入 Dice 项,让网络同时优化像素准确性和区域重叠度。
def dice_loss(pred_logits, target, smooth=1.0): pred = torch.sigmoid(pred_logits) intersection = (pred * target).sum() union = pred.sum() + target.sum() return 1 - (2 * intersection + smooth) / (union + smooth) def combined_loss(pred_logits, target): bce = torch.nn.functional.binary_cross_entropy_with_logits(pred_logits, target) dice = dice_loss(pred_logits, target) return bce + dicesmooth 参数的作用是防止分母为 0,取 1.0 即可。我习惯让 BCE 和 Dice 的权重相同,也就是直接相加。如果发现训练初期 Dice 项不稳定,可以把 Dice 权重降到 0.5,但一般不需要。
评估指标方面,最常用的是 Dice 和 IoU,再配合 AUC 观察分类能力。计算 Dice 时要放在 FOV mask 内,而不是全图。DRIVE 原图有黑色边框,如果全图计算,面积占比不小的边界像素会稀释指标,让结果看起来虚高但不真实。
from sklearn.metrics import roc_auc_score def dice_score(pred_mask, gt_mask, fov_mask): pred = (pred_mask > 0.5).astype(np.uint8) gt = gt_mask.astype(np.uint8) fov = fov_mask.astype(np.uint8) p = pred[fov > 0] g = gt[fov > 0] smooth = 1.0 return (2 * (p * g).sum() + smooth) / (p.sum() + g.sum() + smooth) def auc_score(pred_prob, gt_mask, fov_mask): gt = gt_mask.astype(np.uint8) fov = fov_mask.astype(np.uint8) return roc_auc_score(gt[fov > 0], pred_prob[fov > 0])注意 pred_mask 是二值图,pred_prob 是概率图。AUC 需要的是未经过阈值的概率输出,所以测试时要保留 sigmoid 后的浮点结果。DRIVE 上的常见 Dice 区间是 0.76~0.82,超过 0.82 已经算不错,超过 0.85 需要额外使用大型预训练模型或其他 trick。
4. 训练 UNet:从 DRIVE 训练集到验证集的全过程
4.1 把 patch 封装成 Dataset
在进入训练循环之前,先把前面提取的 patch 封装成 PyTorch Dataset。这样代码更清晰,也方便在训练时随机取样并同步应用数据增强。由于 DRIVE 没有专门的验证集,常见做法是从训练集里留出 2~3 张图作为验证集,剩余 17~18 张做训练。如果你不想留验证集,可以全部训练并只在测试集上看最终结果,但那样就少了调参的依据。
import torch from torch.utils.data import Dataset class DrivePatchDataset(Dataset): def __init__(self, images, labels, fovs, transform=None): self.images = images self.labels = labels self.fovs = fovs self.transform = transform def __len__(self): return len(self.images) def __getitem__(self, idx): img = self.images[idx] lbl = self.labels[idx] fov = self.fovs[idx] if self.transform is not None: transformed = self.transform(image=img, mask=lbl, fov=fov) img = transformed['image'] lbl = transformed['mask'] fov = transformed['fov'] img = torch.from_numpy(img.transpose(2, 0, 1)).float() lbl = torch.from_numpy(lbl).float().unsqueeze(0) fov = torch.from_numpy(fov).float().unsqueeze(0) return img, lbl, fov我在这里把 label 和 fov 都保留了四个维度,方便后续直接与模型输出计算损失。如果你使用的是 CrossEntropyLoss,label 需要改成 LongTensor 并且没有 unsqueeze,但当前是 BCE + Dice,所以保持 float 和单通道即可。数据增强中的 mask 和 fov 参数是让 albumentations 对它们一起做同样的几何变换。
4.2 训练循环与学习率设置
训练 DRIVE 时,我一般设 epoch 为 150,初始学习率 1e-3,batch size 16。优化器选 Adam,weight_decay 给 1e-5 防止过拟合。学习率调度用 ReduceLROnPlateau,当验证集 Dice 连续 10 个 epoch 不改善时,学习率降一半,最低降到 1e-6 就不再降。
model = UNet(in_channels=3, num_classes=1, base=64).to(device) optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-5) scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, mode='min', patience=10, factor=0.5, min_lr=1e-6 ) best_dice = 0.0 for epoch in range(150): model.train() running_loss = 0.0 for img, lbl, fov in train_loader: img, lbl = img.to(device), lbl.to(device) pred = model(img) loss = combined_loss(pred, lbl) optimizer.zero_grad() loss.backward() optimizer.step() running_loss += loss.item() model.eval() valid_dice = evaluate_dice(model, valid_loader, device) scheduler.step(valid_loss) if valid_dice > best_dice: best_dice = valid_dice torch.save(model.state_dict(), 'best_unet_drive.pth') print(f"Epoch {epoch+1:03d} | loss {running_loss/len(train_loader):.4f} | valid dice {valid_dice:.4f}")evaluate_dice 函数的作用是在验证集上把 patch 预测结果重新拼成完整图,然后用 FOV mask 计算 Dice。训练时我通常用 patch 的 label 直接算 Dice,但验证时会把 patch 拼回原图尺寸后再算,这样更接近真实测试环境的指标。注意每次模型切换到 eval 模式后要关闭梯度计算,否则显存会被反向传播的图占满。
这里的 weight_decay 是 Adam 的 L2 正则项。对 DRIVE 这种小数据集,它可以限制模型参数过大的情况,但设得太高会让损失函数收敛变慢,1e-5 是我多次实验后认为最稳的值。如果训练集和验证集 Dice 差距超过 5 个百分点,可以把 weight_decay 加到 1e-4。
4.3 模型保存与断点续训
只保存模型权重虽然方便,但如果你想恢复到某个 epoch 继续训练,就得把 optimizer、scheduler 的状态一起保存。第一版我吃过亏,训练到第 80 个 epoch 时显存崩溃,之前的权重只保存了模型没有保存优化器状态,只能从头跑 80 个 epoch。后来我改成统一保存 checkpoint 文件。
checkpoint = { 'epoch': epoch + 1, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'scheduler_state_dict': scheduler.state_dict(), 'best_dice': best_dice, } torch.save(checkpoint, f'checkpoint_epoch{epoch+1}.pth')恢复训练时要做三件事:加载模型参数、加载优化器参数、把当前 epoch 数设置为你保存的 epoch。如果你用的是 cos 等带周期性的调度器,不恢复 epoch 会导致学习率跳变,前面保存的最佳 Dice 也不会被正确比较。
4.4 测试集评估:输出血管分割图
训练结束后,用测试集做最终评估。测试集 20 张图像的尺寸是 565×584,需要先做和训练相同的 resize 到 576×576,然后预测并 resize 回原始分辨率,再和 ground truth 及 FOV mask 计算指标。
def predict_full_image(model, image_path, device, size=(576, 576)): image = cv2.imread(image_path, cv2.COLOR_BGR2RGB) h, w = image.shape[:2] image = cv2.resize(image, size, interpolation=cv2.INTER_LINEAR).astype(np.float32) / 255.0 img = torch.from_numpy(image.transpose(2, 0, 1)).unsqueeze(0).to(device) model.eval() with torch.no_grad(): pred = torch.sigmoid(model(img)).cpu().numpy()[0, 0] pred = cv2.resize(pred, (w, h), interpolation=cv2.INTER_LINEAR) return pred注意最终 resize 回到原始尺寸时,预测概率图用线性插值不会破坏连续性,但如果你要把它变成二值图,建议先阈值再放大,或者使用最近邻插值,否则边界会出现灰度渐变像素。我一般保存二值图时用pred_mask = (pred > 0.5).astype(np.uint8),并配合原 FOV mask 把黑色背景清掉。
5. DRIVE 训练 UNet 的避坑指南:5 个常见翻车点
5.1 现象:验证集 Dice 一直不变,loss 也不降
训练了 30 个 epoch,验证集 Dice 一直都在 0.2 左右,loss 曲线是一条平线。这种情况最常见的原因是标签和图像在读取或处理时发生了错位,模型看到的是一个与 ground truth 无关的输入,几乎学不到任何有效特征。
原因几乎都出在两个地方:一是使用 cv2.imread 时默认读成 BGR,而数据增强时又把它当成 RGB 处理,导致颜色通道错乱;二是裁剪 patch 时同时修改图像和 mask,但增删操作没有同步,标签被移动了一个偏移量。检查方法很简单,随机选几个 patch,把 image、label 拼在一起用 matplotlib 可视化,肉眼看血管和标注是否在同一个位置。
解决方法是写一个 debug 脚本,单独把每一个 Dataset 输出可视化出来,确认没有错位后再进入训练。这个步骤虽然麻烦,但能省掉后面排查异常指标的大把时间。
5.2 现象:预测图全黑或全白
测试集上预测结果几乎全是背景,或者全是血管,Dice 在 0.05 以下。全黑通常是模型学会了把所有像素都预测为背景,全白则相反。这两个问题看似相反,但根源往往是类别不平衡处理不到位。
血管占比小,普通 BCE 会在初始阶段把预测推向背景,如果不加入 Dice 项或者不使用加权采样,模型会停留在全背景的局部最优。全白我遇到过两次,一次是因为标签在读取后没有二值化,导致 ground truth 里同时有 0、127、255 三种值,模型被中间灰度值误导;另一次是损失函数里 Dice 项写错,把目标当成了预测的相反值。
解决方法是先检查 label 的唯一取值,确认只有 0 和 1,再检查 loss 计算中 pred 和 target 的通道是否匹配。最后用单个 batch 做快速测试:在训练时打印出 pred 的均值,如果全部小于 0.1,就要考虑提高正样本权重。
5.3 现象:训练集 Dice 高但验证集 Dice 低
训练集 Dice 到 0.92,验证集只有 0.70,明显过拟合。DRIVE 训练样本太少,20 张图里如果留出 3 张做验证,训练只有 17 张,UNet 参数量大,很容易把训练图背下来。
解决过拟合有几条路。第一是加强数据增强,特别是弹性形变和亮度对比度变化,让模型不能简单依赖记忆。第二是把 weight_decay 调到 1e-4 或 1e-3,限制参数规模。第三是减少 base 通道数,从 64 降到 32,UNet 的通道数减半后参数量大幅下降,对 DRIVE 这类小数据反而效果更好。我实际对比后,base=32 比 base=64 的验证 Dice 高 1~2 个点,因为模型容量小,更难过拟合。
5.4 现象:血管中心粗大但末梢断裂
预测结果看起来像“水晶管”,主干血管很粗,毛细血管断成一截一截。这个现象说明模型只保留到了高层语义信息,丢失了低层边缘细节,跳跃连接没有起到应有的作用。
原因可能是输入分辨率太低,比如 resize 到 256 后细血管直接消失;也可能是数据增强里的弹性形变强度太大,把血管末梢拉断了。另一个容易被忽略的原因是 patch 尺寸太小,血管末梢在整个 patch 中占比极低,模型认为它是可以忽略的背景。
解决方法是把输入分辨率恢复到 576,弹性形变 alpha 降到 5 以内,同时训练时考虑让裁剪窗口完全覆盖眼底中心区域附近,这样模型有更多机会看到密集的小血管。推理阶段还可以用重叠滑动窗口,避免 patch 边缘血管被裁掉。
5.5 现象:显存不够,训练直接 OOM
使用 576×576 整图训练,batch size 为 8 时显存溢出。DRIVE 图像本身不大,但 UNet 的中间特征图在编码器第一层就有 64 通道,后续翻倍,显存消耗并不低。
解决思路有几个:一是把训练输入切成 patch,这是最有效的降显存手段;二是降低 batch size,这会影响 BN 统计量,必要时应把 BatchNorm 换成 InstanceNorm;三是关闭不需要的梯度,比如在对验证集做评估时用with torch.no_grad();四是使用混合精度训练,PyTorch 的自动混合精度能减少一半显存占用。推荐优先做第一点,因为对分割效果影响最小。
6. 把 Dice 从 0.78 拉到 0.82:三个不改变模型结构的技巧
模型训练稳定之后,如果你想让 Dice 再往上走一步,不需要换更大的网络,下面三个技巧在 DRIVE 上效果明显。
第一个是重叠滑动窗口推理。训练时如果用 patch,预测时直接整图输入可能出现细节差异,因为模型见过的最大尺寸是 patch,整图上采样时感受野变化不大,但边界位置容易失真。我在测试时把每张图切成 256×256 的窗口,步长设为 128,每个像素会出现在多个窗口的预测结果中,最后取平均值作为该像素的概率。这个操作几乎白捡 1~2 个点的 Dice,且不改变模型权重。
第二个是测试时增强。对每张测试图做水平翻转、垂直翻转和旋转 90 度这几种变换,分别得到概率图后反变换回原方向再平均。TTA 对血管分割很友好,因为血管走向各异,多视角预测能互相补充断裂的毛细血管。加上 TTA 后 Dice 通常能提升 0.5 到 1 个点,代价只是推理时间变成原来的 4~8 倍。
第三个是预测后处理。DRIVE 的 ground truth 是连续的血管树,但模型预测出的概率图往往会在细小分支处出现不希望的孤立点。我一般用scipy.ndimage.median_filter对概率图做一个 2×2 的中值滤波,再去掉连通域小于 10 像素的区域。不过注意后处理不能太激进,否则会把薄血管一起抹掉。DRIVE 上我对比过,去掉小连通域后 Dice 上升了 0.3 个点,但灵敏度有轻微下降。
我之前在模拟项目X里做这个任务时,卡在 Dice 0.78 很久,后来发现只是测试时直接整图预测,patch 训练和整图推理之间的尺度不一致让边界血管预测得很差。改成重叠滑动窗口后直接到了 0.80,再叠加 TTA 和轻量后处理,最终停在 0.82。这三个技巧有一个共同点:都不动模型结构,只改变预测时的数据流转,所以它们很容易迁移到其他分割任务上。
如果以后你换到更大的医学图像数据集,这套流程依然适用:先用 patch 训练解决显存和样本量问题,再用重叠窗口 + TTA 去掉推理噪声,最后针对目标器官做轻量后处理。血泪教训是不要在调模型结构上花太多时间,先把数据管线和推理流程做扎实,Dice 往往自己就涨上去了。希望帮到你。
本文还有配套的精品资源,点击获取