☰
基于UNet的视网膜血管分割实战:DRIVE数据集与训练闭环
2026/9/26 8:02:53 网站建设 项目流程

简介:基于UNet架构的视网膜血管分割项目,使用PyTorch框架,并采用DRIVE公开数据集进行模型训练与测试,是一份完整可复现的医学图像分割方案。资源主要面向深度学习初学者、医学影像研究人员以及需要血管分割基准实验的开发者,能够解决从原始视网膜图像到血管结构提取的端到端流程搭建问题。项目覆盖数据预处理脚本、模型定义、训练测试流程与可视化工具,预处理步骤包括图像标准化、增强、去噪和对比度调整,有助于提高血管结构的可视性,帮助读者理解UNet在细粒度医学结构分割中的实际应用与调优思路。包内共34个文件,以Python脚本、PNG结果图为主,兼有数据集压缩包、依赖清单、说明文档和附录资料,整体约36.81MB,目录结构清晰,功能模块划分明确。目前已有133人学习下载,适合复现实验、二次开发或作为医学图像处理课程设计参考。

1. 基于 UNet 架构的视网膜血管分割:从 DRIVE 数据集到训练闭环

第一次把深度学习实战押在视网膜血管分割上,是个看起来冷门、实际很能练手的决定。基于 UNet 架构的视网膜血管分割项目把 DRIVE 公开数据集、PyTorch 实现、数据预处理、训练和可视化全串在同一条流水线里,你拿到的不再是一个孤零零的模型文件,而是从原始眼底图到最终分割结果的完整工程。血管是典型的细长目标,比通用语义分割更敏感,对比度、裁剪方式、掩码处理都会让最终指标产生明显波动,这反倒让它成为理解分割任务弊端的绝佳样例。想跑通 UNet 或准备进入医学图像方向的初学者,都能靠这套代码省掉不少时间。你只需要配好 PyTorch 环境,顺着脚本往下执行,就能看到训练曲线和血管预测图一步步稳定下来。

这个项目也解决了一个常见尴尬:很多人照着经典论文写模型很快,但栽在数据处理和评估细节上。DRIVE 数据集的训练集、测试集、掩码目录各有各的坑,如果无人指路,第一个 epoch 就可能出现“损失在降、指标不动”的怪现象。后面我会把这些地方的错误表现和排查方法全部拆开讲,也会把参数设置的依据说清楚,方便你替换到自己的数据集上继续改进。

2. 分层拆解 UNet:血管分割选它,到底选在哪

2.1 编码器-解码器结构与跳跃连接的实际含义

UNet 的结构看起来对称,但它真正厉害之处是把“上下文”和“细节”两条信息通路同时保留下来。左边的编码器逐步下采样,每一层都在扩大感受野,网络能知道一根血管是在视网膜中央还是边缘、周围有没有其他组织干扰。右边的解码器再把特征图一步步恢复回原分辨率,为每个像素给出最终判断。

只靠“压缩-恢复”还不够,因为血管边缘很细,下采样次数一多,细小的分支信息就会在池化过程中丢掉。跳跃连接就是专门来补这个缺陷的:把编码器同尺度特征图横向拼到解码器上,浅层特征提供精确位置,深层特征提供语义判断。这也是 UNet 在医学图像分割里成为默认基线的核心理由——它的结构让“细”和“宽”都能保住。

在实际操作这套项目时,你会看到它在每个编码块里用了两次卷积加 ReLU,通道数从 64 开始逐层翻倍。这是非常经典的设计,不是拍脑袋定的。首层通道过少,模型对血管边缘的拟合能力不足;首层通道过多,DRIVE 这种几千张级别的小数据集很容易过拟合。64 个初始通道基本是兼顾参数量和表达力的经验值,如果换到更大规模数据集,可以考虑把初始通道调到 128。

2.2 损失函数与评估指标:从 BCE 到 Dice 的组合逻辑

血管分割是一个二分类问题,每个像素非血管即血管,但正负样本比例极度失衡。DRIVE 数据集里血管像素通常只占百分之十几,其余都是背景。如果直接用普通交叉熵,模型会把几乎所有像素预测成背景,损失值看起来不高,实际分割结果就是一片黑。

这套项目里比较合理的做法是把 BCE 和 Dice Loss 组合在一起使用。BCE 保留像素级的梯度信号,Dice Loss 则直接优化区域重叠度,对正负样本不均衡不那么敏感。你也可以两手都用:前几十个 epoch 让 BCE 主导学习基础特征,后面再把 Dice 权重调高来细化边界。这个策略很少有人直接写进文档,但跑起来你会发现它对收敛稳定性帮助很大。

评估指标也不要只盯 IOU(交并比)。在医学图像里,Dice Coefficient、敏感度(Sensitivity)和特异性(Specificity)往往更能说明问题。敏感度低了,说明血管被漏掉太多,这对临床诊断是致命的。所以看训练日志时,不要只看总 loss,要额外打印这些指标。项目自带的训练日志和可视化工具通常都会输出这几个数值,如果你拿到代码后发现没打印,建议自己在验证集上补一段评估逻辑,避免训练完只得到一个“看起来不错”的模型。

2.3 UNet 的 PyTorch 实现代码与参数说明

先看网络主体。下面这段代码是这种项目里最常见的写法,省略了部分重复块,保留了完整结构逻辑,方便你对照源码理解。

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, out_channels=1, init_features=64): super().__init__() features = init_features # 编码器 self.enc1 = DoubleConv(in_channels, features) self.enc2 = DoubleConv(features, features * 2) self.enc3 = DoubleConv(features * 2, features * 4) self.enc4 = DoubleConv(features * 4, features * 8) self.pool = nn.MaxPool2d(kernel_size=2, stride=2) # 瓶颈 self.bottleneck = DoubleConv(features * 8, features * 16) # 解码器 self.up4 = nn.ConvTranspose2d(features * 16, features * 8, kernel_size=2, stride=2) self.dec4 = DoubleConv(features * 16, features * 8) self.up3 = nn.ConvTranspose2d(features * 8, features * 4, kernel_size=2, stride=2) self.dec3 = DoubleConv(features * 8, features * 4) self.up2 = nn.ConvTranspose2d(features * 4, features * 2, kernel_size=2, stride=2) self.dec2 = DoubleConv(features * 4, features * 2) self.up1 = nn.ConvTranspose2d(features * 2, features, kernel_size=2, stride=2) self.dec1 = DoubleConv(features * 2, features) self.out_conv = nn.Conv2d(features, out_channels, kernel_size=1) def forward(self, x): # 编码路径 enc1 = self.enc1(x) enc2 = self.enc2(self.pool(enc1)) enc3 = self.enc3(self.pool(enc2)) enc4 = self.enc4(self.pool(enc3)) # 瓶颈 bottleneck = self.bottleneck(self.pool(enc4)) # 解码路径 dec4 = self.up4(bottleneck) dec4 = torch.cat([dec4, enc4], dim=1) dec4 = self.dec4(dec4) dec3 = self.up3(dec4) dec3 = torch.cat([dec3, enc3], dim=1) dec3 = self.dec3(dec3) dec2 = self.up2(dec3) dec2 = torch.cat([dec2, enc2], dim=1) dec2 = self.dec2(dec2) dec1 = self.up1(dec2) dec1 = torch.cat([dec1, enc1], dim=1) dec1 = self.dec1(dec1) return self.out_conv(dec1)

这段代码里有几个需要留意的点。解码器每层上采样之后都要和编码器对应层的输出做通道拼接,torch.cat的维度是dim=1,也就是通道维。很多新手在这里拼错维度,或者忘记拼,直接导致特征融合失效,训练结果表现很差。

ConvTranspose2d负责上采样,但它不是唯一选择。项目里如果用双线性插值配合卷积,效果通常也差不多,但参数量和训练表现略有差异。实践中替换成nn.Upsample(scale_factor=2, mode='bilinear')后,显存占用会低一些,适合显存吃紧的机器。两种做法都可以保留,不要盲目改。

in_channels默认是 3,对应 RGB 三通道;如果按后面的预处理方案只保留绿色通道,这里要改成 1。out_channels是 1,表示最终输出单通道预测图,再接Sigmoid得到血管概率图。

3. 数据预处理与加载:DRIVE 的目录、掩码和分布防线

3.1 DRIVE 数据集结构说明

DRIVE 是视网膜血管分割最常用的公开数据集之一,全部来自糖尿病视网膜病变筛查项目。标准目录结构通常是这样的:

目录/文件内容说明
training/images/训练集眼底图(.tif)通常 20 张,包含原始 RGB 图
training/1st_manual/训练集人工标注(.gif)血管标注,白线表示血管
training/mask/训练集 ROI 掩码(.gif)标记眼底视网膜有效区域
test/images/测试集眼底图(.tif)也是 20 张,不参与训练
test/1st_manual/测试集人工标注(.gif)用于最终评估
test/mask/测试集 ROI 掩码(.gif)评估时只统计掩码内部区域

拿到项目后,第一件事不是跑训练,而是先确认脚本是否把mask和1st_manual区分开。我见过不少复现失败的项目,就是把掩码当成了标签去算损失,结果模型学到的根本不是血管,而是眼底图像的外轮廓。DRIVE 的mask标记的是“哪些区域要参与评估”,1st_manual才是血管金标准,这两个文件一旦读混,后面所有指标都会失真。

另一个容易被忽略的地方是测试集也有自己的掩码。评估时要把预测结果乘上掩码,只统计视网膜有效区域,否则图片黑色边框会被算进背景里,虚高特异性分数。顺手把测试集掩码读进来,训练最后阶段评估时用上,是省事又稳妥的习惯。

3.2 预处理脚本:绿色通道、CLAHE 与归一化

眼底图里血管在绿色通道下对比度最高,红色通道偏亮容易饱和,蓝色通道噪声大。常见的预处理做法是分离通道后只保留绿色通道,再做对比度受限自适应直方图均衡化(CLAHE),最后归一化。这样既压缩了计算量,又让血管边缘更清晰。

下面这段预处理逻辑和项目里常见实现基本一致:

import cv2 import numpy as np from glob import glob import os def preprocess_drive_images(image_dir, save_dir): os.makedirs(save_dir, exist_ok=True) image_paths = sorted(glob(os.path.join(image_dir, "*.tif"))) for path in image_paths: img = cv2.imread(path) # 分离 BGR 通道,保留绿色通道 b, g, r = cv2.split(img) # CLAHE 提升局部对比度,clipLimit 控制增强幅度 clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8)) g_clahe = clahe.apply(g) # 归一化到 [0,1],后面转 Tensor 时再换算成 [0,1] float g_norm = g_clahe / 255.0 base_name = os.path.basename(path).replace(".tif", ".npy") np.save(os.path.join(save_dir, base_name), g_norm.astype(np.float32))

绿色通道加 CLAHE 是眼底图像分割里非常固定的组合拳,clipLimit取 2.0 是比较稳的值。如果发现血管和背景对比还是不够,可以适当把clipLimit调到 3.0,但别调太高,否则背景噪声也会被放大。这里还有一个容易暗藏的问题:cv2.imread读进来的是 BGR 顺序,不是 RGB,如果你在 NumPy 里直接按索引取r通道,取到的实际上是红色通道。务必记住 OpenCV 通道顺序这个细节。

保存成.npy而不是图片格式,看起来多此一举,实际有两个好处:一是省去训练时反复读取磁盘和 JPEG 解压的时间;二是numpy格式可以直接进Dataset,省去一次transform字符串解析。整套预处理只需要跑一次,之后训练和验证都用同一份预处理产物,也避免了训练时做在线 CLAHE 导致 CPU 占用过高的问题。

3.3 数据增强与 Dataset 加载器设计

训练集只有 20 张图,如果不做增强,UNet 这种大参数模型很快就会过拟合。使用随机裁剪、水平翻转、垂直翻转三种基础增强足够,因为血管本身对翻转不敏感。额外的旋转或多尺度裁剪也可以加,但要注意掩码必须和图像应用同一种变换。这个“同步变换”的细节在增强代码里最容易出错。

下面是一份标准 Dataset 结构,图像和掩码各自经过相同随机种子下的同步变换:

import torch from torch.utils.data import Dataset import numpy as np import random class DRIVEDataset(Dataset): def __init__(self, image_dir, mask_dir, crop_size=256, augment=False): self.image_paths = sorted(glob(os.path.join(image_dir, "*.npy"))) self.mask_paths = sorted(glob(os.path.join(mask_dir, "*.gif"))) self.crop_size = crop_size self.augment = augment def __len__(self): return len(self.image_paths) def __getitem__(self, idx): # 这里实际项目会按文件名对应读取,而不是按索引硬匹配 image = np.load(self.image_paths[idx]) mask = cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE) mask = (mask > 0).astype(np.float32) # 随机裁剪到固定尺寸,训练和验证都用同一套逻辑 H, W = image.shape x = random.randint(0, H - self.crop_size) y = random.randint(0, W - self.crop_size) image = image[x:x + self.crop_size, y:y + self.crop_size] mask = mask[x:x + self.crop_size, y:y + self.crop_size] if self.augment: if random.random() > 0.5: image = np.flip(image, axis=1) mask = np.flip(mask, axis=1) if random.random() > 0.5: image = np.flip(image, axis=0) mask = np.flip(mask, axis=0) image_tensor = torch.from_numpy(image.copy()).unsqueeze(0) mask_tensor = torch.from_numpy(mask.copy()).unsqueeze(0) return image_tensor, mask_tensor

这份代码的裁剪逻辑是固定的正方形区域,如果原始图和掩码尺寸不一致,必须先统一尺寸。DRIVE 的原始图像是 565×584,通常项目会先把图和掩码 resize 到统一大小,再执行裁剪。这里需要特别关注mask的读取方式:.gif是单通道,直接用IMREAD_GRAYSCALE读;如果你用 PIL 读,也要显式转成 L 模式,避免出现通道维度错位。

torch.from_numpy(...).unsqueeze(0)的作用是给二维数组加上一个通道维。如果你的网络输入是单通道,这里正好匹配;如果仍然想用三通道输入,保留绿色通道的同时也可以把红蓝通道作为弱特征一起输入,就要在预处理阶段合并成三通道数组。两种方式都有人用,实验下来绿色通道单输入在 DRIVE 上表现并不差,而且训练更快。

4. 训练与调参落地:优化器、损失和检查点的工程化写法

4.1 训练主循环:从模型初始化到 checkpoint 保存

搭建完网络和 Dataset,训练环节实际上只剩下一个标准循环。这个项目里的训练逻辑通常包含训练阶段评估、验证阶段评估、模型保存三部分。下面这段伪代码可以从源码中对应出来:

model = UNet(in_channels=1, out_channels=1).cuda() optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-5) criterion = build_bce_dice_loss() for epoch in range(epochs): model.train() train_loss = 0.0 for images, masks in train_loader: images = images.cuda() masks = masks.cuda() preds = model(images) loss = criterion(preds, masks) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() train_loss += loss.item() # 每 10 个 epoch 保存一次完整快照 if epoch % 10 == 0 or epoch == epochs - 1: torch.save({ "epoch": epoch, "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "train_loss": train_loss / len(train_loader), }, f"checkpoints/unet_epoch_{epoch}.pth")

以上代码的优化器选的是 AdamW,它和经典 Adam 的区别是权重衰减处理方式更规范,在小数据集上更容易得到稳定结果。clip_grad_norm_是一个经常被拿掉的步骤,但血管分割数据集数值异常多,偶尔会出现个别样本产生巨大梯度。加一行裁剪,最多损失一点收敛速度,却能避免训练中期 loss 突然变成 NaN。

checkpoint 保存不能只存模型权重。把 epoch、优化器状态和当前 loss 一起存下来,后面想恢复训练或者调整学习率继续跑,都能直接通过torch.load恢复。若是只存model.state_dict(),中断恢复后优化器状态会重新初始化,相当于变相丢掉了之前的学习率调整记录。

4.2 关键参数配置与调整基准

下面这张参数表是从 DRIVE 上跑通这类项目的常见配置中整理出来的。不同脚本会有出入,但数量级基本一致。

参数推荐值调整说明
输入尺寸256×256显存小就降到 224,但标签边缘会受影响
Batch Size420 张训练图,batch 再大容易过拟合
初始学习率1e-4用 1e-3 容易前几十个 epoch 震荡
学习率调整ReduceLROnPlateau验证 Dice 不升时降 0.5 倍
Epoch100-150DRIVE 上 100 epoch 基本足够
优化器AdamW权重衰减建议 1e-5 到 1e-4
损失组合BCE + Dice两权重建议各 0.5

需要注意,batch_size=4在单卡 8GB 显存下跑 256×256 输入通常没有问题。如果显存紧张,可以优先缩小输入尺寸而不是降低 batch size,因为 UNet 的参数量摆在那里,批大小过小会导致 BatchNorm 统计不稳定,训练振荡严重。

学习率 1e-4 是我在这个项目上的常用起点。上了 AdamW 也不要迷信默认学习率,尤其当输入尺寸变小后,梯度规模变化,默认值不一定合适。验证集 Dice 连续 15 个 epoch 不提升,就把学习率乘以 0.5;连续 30 个 epoch 不动,就要检查数据预处处理是不是出了问题。

4.3 可视化工具:把预测结果和损失曲线拉出来看

该项目里自带的可视化工具一般做三件事:绘制训练损失曲线、保存预测概率图、把预测结果和金标准叠加对比。很多跑深度学习项目的习惯是把可视化当成事后工作,但我建议每个 epoch 结束后都保存一次小图。20 张训练图的模型,看不出来 loss 数值是否稳定,但一定看得出来分割边缘是否成型。

可视化代码的常见形式如下:

def visualize_prediction(image, mask, pred, save_path): import matplotlib.pyplot as plt plt.figure(figsize=(12, 4)) plt.subplot(1, 3, 1) plt.imshow(image.squeeze(), cmap="gray") plt.title("Input") plt.subplot(1, 3, 2) plt.imshow(mask.squeeze(), cmap="gray") plt.title("Ground Truth") plt.subplot(1, 3, 3) plt.imshow(pred.squeeze(), cmap="gray") plt.title("Prediction") plt.savefig(save_path, dpi=150, bbox_inches="tight") plt.close()

保存的预测图应该是模型输出的概率图,是 0 到 1 之间的浮点数,不要提前用阈值二值化,因为你看灰度图能直接判断概率分布的高斯性和边界清晰度。二值化阈值通常会设在 0.5,这在血管分割里不是最优选择,后文验证部分会提到。

如果发现预测图上血管整体比标签细一圈,问题大概率出在损失函数里 Dice 权重过低;如果预测图有许多零散噪点,大概率是数据增强不够或者模型过拟合。可视化结果的价值就在这里:不用等训练到最后一个 epoch,你就能提前判断方向是否正确。

5. 避坑指南:DRIVE 数据集和 UNet 训练最容易翻车的五处

5.1 掩码被当成标签读进损失函数

  • 现象:训练第一轮 loss 很低,验证集 Dice 却几乎为 0。
  • 原因:把mask目录下的 ROI 掩码当作血管金标准。ROI 掩码标注的是整个视网膜圆形区域,和血管分割标签长相完全不同,模型很容易拟合出圆形轮廓。
  • 解决:先打印一个 batch 的 mask 张量,看一眼标签是不是只有血管纹理;再核对Dataset里读取路径指向的是1st_manual而不是mask。这种低级错误最隐蔽,因为日志中 loss 完全正常。

5.2 随机裁剪把重要血管分支切掉

  • 现象:用 256×256 随机裁剪训练 100 个 epoch 后,测试集表现比训练集差不少,而且预测图边缘区域断裂严重。
  • 原因:DRIVE 眼底图是圆形 ROI,边缘本身有大量黑色背景。随机裁剪会频繁采到背景区域,同时血管分支在原始图中贯穿全图,固定裁剪尺寸把长程连接切断。
  • 解决:训练时把裁剪中心限制在 ROI 掩码内部,或者先做一次基于掩码的 ROI 外接矩形裁剪,再在矩形内部随机裁剪。另一个辅助办法是增强时加入小幅旋转和弹性形变,让血管分支走向更多样,而不是依赖裁剪保留连通性。

5.3 验证时候图片尺寸和模型输入不一致

  • 现象:训练时 loss 正常,到验证阶段突然报维度错误。
  • 原因:验证阶段没有随机裁剪,直接输入原始尺寸 565×584,而 UNet 下采样四次后要求输入尺寸能被 16 整除。
  • 解决:统一验证流程,所有输入先resize到 256×256,或者用F.interpolate把输出插值回原图尺寸。不能直接让模型吃原图,因为 UNet 结构对非整除尺寸很敏感,最后一层上采样的尺寸会错位。

5.4 评估指标忽略测试集掩码

  • 现象:测试集 Dice 看起来很漂亮,但打开预测图发现黑色边框全被预测成背景。
  • 原因:背景区域占比例很大,模型把黑色区域预测为背景能刷高特异性,而边界区域没有参与真实评估。
  • 解决:计算指标前,把预测结果和 ROI 掩码做逐像素相乘,同时把金标准和预测结果都限定在掩码区域。正确公式是用掩码内的预测去算 Dice 和敏感度,而不是用全图。

5.5 显存不足时盲目调低 batch size

  • 现象:显存报错后把 batch size 从 4 改到 1,训练开始震荡。
  • 原因:BatchNorm 在 batch size 为 1 时统计量是单样本均值,方差失真,相当于动态噪声,导致模型无法稳定收敛。
  • 解决:优先把输入尺寸从 256 降到 224,或者关闭 BatchNorm 并改用 InstanceNorm。代码上如果把nn.BatchNorm2d替换成nn.InstanceNorm2d,batch size 为 1 时依然能稳定训练。这个替换对单张预测也友好,很多和本项目类似的代码默认用 BatchNorm,你得在遇到显存瓶颈时想起来这层关系。

5.6 保存预测图和保存概率图混淆

  • 现象:测试集可视化的血管都纤细且星点状,和金标准差距大。
  • 原因:直接对 sigmoid 输出做> 0.5的硬阈值,但血管概率分布往往不是标准的二值分布,在 0.5 附近有大量像素。
  • 解决:先把预测图保存成概率图观察灰度分布,取值区间代表血管的置信度;如果需要二值化,先做 Otsu 阈值或者验证集上搜索最优阈值。

6. 验证与改进:跑完测试集,再做一次细节复盘

6.1 用一次系统化预测来检验整个流程

训练停下后,不要急着调参,先完整跑一遍测试集预测脚本。很多项目把测试集预测和训练分开,你可以把每张测试图的输入、ROI 掩码、预测概率图和二值化结果都存档,然后按掩码区域计算 Dice 和敏感度。记录这三个数值的同时,也保存概率图的灰度分布统计,比如 0.1 以下的像素占比多少、0.5 到 0.9 之间的像素占比多少。这些统计能反映出模型是犹豫型还是极端型,在血管分割里,过度极端的输出往往代表边界信息丢失。

测试集只有 20 张图,逐张看预测图是说得过去的。找到预测最差的几张图,把原始图像、绿色通道增强图、金标准和预测四个图放在一起对照。通常你会发现两类问题:一是图像本身对比度差,预处理参数不适合这张图;二是细小血管在深层特征里被稀释,深度学习模型对细长分支的还原天然有瓶颈。

6.2 指数移动平均、测度与阈值搜索

在推进到第 6 章时,你已经有了一个可以稳定收敛的模型基线。下面这个技巧能直接把测试集 Dice 提高一到两个点。传统模型预测时直接用当前权重,但训练后期权重在最优解附近震荡。改用指数移动平均(EMA)来保存一份滑动平均权重,往往能得到更平滑的预测结果。PyTorch 里的实现不复杂,核心代码长这样:

ema_decay = 0.99 ema_model = {k: v.clone() for k, v in model.state_dict().items()} # 每个 step 更新一次 for k in model.state_dict(): ema_model[k] = ema_decay * ema_model[k] + (1 - ema_decay) * model.state_dict()[k]

EMA 在分类任务里能稳定验证集分数,在血管分割里同样适用。另外,评估阈值不要死守 0.5。在验证集上画出 Dice 随阈值变化的曲线,最优阈值往往会偏到 0.3 到 0.6 之间某一处。这个“搜索阈值”的步骤非常短,对最终指标的影响却很明显,是复现高分报告时不可跳过的细节。

6.3 给后续实验留下扩展接口

基线跑通后,如果你想基于 UNet 做改进,最优改动不是换网络结构,而是在数据侧做文章。把绿色通道和红色通道的差值做成额外输入通道,能增强血管和背景的可区分度;或者在预处理阶段加入血管增强滤波器,比如 Frangi 滤波器,它的响应图可以作为第四通道输入。项目里如果预留了通道拼接逻辑,你只需要增加预处理输出和in_channels参数。

很多从公开项目复现的人拿到代码后容易陷入“只训练、不改代码”的状态。我会在每次训练完成后,把测试集预测最差的三张图单独建一个文件夹,记录它们的文件名、Dice 值、预测概率均值。下一次实验结束后,先看这三个文件的指标有没有提升,而不是看平均 Dice。平均 Dice 有时接近了,但最差样本仍然在拖后腿,这在实际应用里恰恰是关键问题。

我从开始完整跑通这个基于 UNet 架构的视网膜血管分割项目起,就养成了一个强制习惯:每次训练落地前,先手动验证 3 张测试图的预处理产物,确认掩码读取正确、增强同步正常,再启动训练。那次把 mask 当标签读取导致整个实验翻车后,这个检查动作就成了我再也跳不过去的流程。希望这个习惯和这套项目的完整拆解,也能让你在 DRIVE 和后续自己的数据上少走一段弯路,希望帮到你。

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

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

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

立即咨询