PyTorch UNet视网膜血管分割实战:从DRIVE数据集到0.85 Dice

发布时间:2026/10/11 22:55:01
PyTorch UNet视网膜血管分割实战:从DRIVE数据集到0.85 Dice 简介这份资源面向医学图像处理方向的深度学习学习者与研究人员提供基于UNet架构的视网膜血管分割完整项目采用PyTorch框架实现从数据预处理、模型训练到测试评估的全流程适合具备一定深度学习基础、希望上手医学图像分割实战的开发者参考。压缩包共34个文件约36.81MB以20张png结果图、7个py脚本为主另含txt说明、docx附赠文档、license及md说明等脚本涵盖模型定义、数据集加载、损失函数、预处理与训练测试等模块结构清晰便于按需查阅。目前已有135人学习下载。项目选用DRIVE公开数据集进行训练与测试配套数据预处理脚本可完成标准化、增强与去噪等操作可视化工具则便于直观比对分割结果并分析血管结构提取效果为视网膜疾病的早期诊断研究提供了一套可复现、易上手的实践方案。1. 视网膜血管分割这套 UNet 方案到底能解决什么临床级问题眼底照相机拍出来的图医生看的是血管的粗细、走向和分支形态糖尿病视网膜病变、青光眼、高血压视网膜病变的判断都绕不开血管的精确提取。但手工勾血管这件事一张 512×512 的 DRIVE 图像熟手也要 20 分钟以上而且不同医生勾出来的边界能差出好几个像素。基于 UNet 架构的视网膜血管分割项目要解决的就是把这个过程自动化输入一张彩色眼底图输出一张二值血管掩膜像素级判断每个位置是不是血管。这套方案适合三类人刚学完 PyTorch 基础、想找一个完整深度学习流程练手的工程师做医学图像处理、需要快速搭一个血管分割 baseline 的研究生以及想把眼底筛查往自动化方向推的产品团队。DRIVE 公开数据集只有 40 张图训练集 20 张、测试集 20 张量小但标注质量高是视网膜血管分割领域最经典的入门基准。用 PyTorch 实现 UNet 做这个任务代码量不大但数据预处理的坑、损失函数的选型、评估指标的解读每一个都能让你卡上半天。下面按「先跑通、再调优、最后避坑」的顺序把整套流程拆开讲。2. 用 PyTorch 搭 UNet从 DRIVE 数据加载到前向传播2.1 DRIVE 数据集的目录结构与预处理脚本DRIVE 原始数据下载下来后目录结构通常是这样的DRIVE/ ├── training/ │ ├── images/ # 20 张训练原图格式 .tif │ ├── 1st_manual/ # 20 张专家手工标注血管掩膜 │ └── mask/ # 20 张 FOV 有效区域掩膜 └── test/ ├── images/ # 20 张测试原图 ├── 1st_manual/ # 测试集标注用于评估 └── mask/ # 测试集 FOV 掩膜这里有个容易翻车的点DRIVE 的标注图是灰度图血管像素值不是 255 而是 1直接当二值图用会出问题。预处理脚本必须做归一化。我一般会写一个DRIVEDataset类继承torch.utils.data.Dataset把图像和掩膜同步做增强。import os import torch from torch.utils.data import Dataset from PIL import Image import numpy as np import torchvision.transforms.functional as TF import random class DRIVEDataset(Dataset): def __init__(self, root_dir, splittraining, transformTrue): self.img_dir os.path.join(root_dir, split, images) self.mask_dir os.path.join(root_dir, split, 1st_manual) self.fov_dir os.path.join(root_dir, split, mask) self.transform transform # 只取 .tif 文件避免读到系统隐藏文件 self.ids sorted([f for f in os.listdir(self.img_dir) if f.endswith(.tif)]) def __len__(self): return len(self.ids) def __getitem__(self, idx): name self.ids[idx] img Image.open(os.path.join(self.img_dir, name)).convert(RGB) mask_name name.replace(.tif, _manual1.gif) # 标注是 gif 格式 mask Image.open(os.path.join(self.mask_dir, mask_name)).convert(L) fov Image.open(os.path.join(self.fov_dir, name.replace(.tif, _mask.gif))).convert(L) img np.array(img, dtypenp.float32) / 255.0 mask np.array(mask, dtypenp.float32) mask (mask 0).astype(np.float32) # 关键把 1 和 0 统一成 0/1 fov (np.array(fov) 0).astype(np.float32) img torch.from_numpy(img).permute(2, 0, 1) # HWC - CHW mask torch.from_numpy(mask).unsqueeze(0) fov torch.from_numpy(fov).unsqueeze(0) if self.transform: # 同步随机翻转图像和掩膜必须用同一组参数 if random.random() 0.5: img TF.hflip(img) mask TF.hflip(mask) fov TF.hflip(fov) if random.random() 0.5: img TF.vflip(img) mask TF.vflip(mask) fov TF.vflip(fov) return img, mask, fov这段代码里最值得说的是mask (mask 0).astype(np.float32)这一行。DRIVE 的标注图在灰度模式下血管区域像素值是 1背景是 0但经过 PIL 读取和 numpy 转换后有些版本会变成 255 和 0。不做二值化统一后面算损失的时候会出现梯度爆炸。fov掩膜的作用是标记圆形视野的有效区域计算损失和指标时只在这个区域内算否则黑色边框会被当成背景拉高准确率但实际分割效果很差。2.2 UNet 网络结构的 PyTorch 实现与通道数配置UNet 的结构不复杂编码器四次下采样解码器四次上采样中间用跳跃连接把编码器的特征拼到解码器对应层。但有几个参数必须根据 DRIVE 的图像尺寸来定。import torch.nn as nn import torch.nn.functional as F class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() 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, x): return self.conv(x) class UNet(nn.Module): def __init__(self, in_ch3, out_ch1, features[64, 128, 256, 512]): super().__init__() self.downs nn.ModuleList() self.ups nn.ModuleList() self.pool nn.MaxPool2d(2, 2) # 编码器 for f in features: self.downs.append(DoubleConv(in_ch, f)) in_ch f # 解码器 for f in reversed(features): self.ups.append(nn.ConvTranspose2d(f*2, f, 2, 2)) self.ups.append(DoubleConv(f*2, f)) self.bottleneck DoubleConv(features[-1], features[-1]*2) self.final nn.Conv2d(features[0], out_ch, 1) def forward(self, x): skip [] for down in self.downs: x down(x) skip.append(x) x self.pool(x) x self.bottleneck(x) skip skip[::-1] for i in range(0, len(self.ups), 2): x self.ups[i](x) s skip[i//2] # 如果尺寸不匹配就裁剪DRIVE 图像 512x512 一般不会出问题 if x.shape ! s.shape: x F.interpolate(x, sizes.shape[2:], modebilinear, align_cornersTrue) x torch.cat([s, x], dim1) x self.ups[i1](x) return torch.sigmoid(self.final(x))features[64, 128, 256, 512]是标准配置显存不够就砍到[32, 64, 128, 256]分割精度会掉 1 到 2 个 Dice 点。ConvTranspose2d做上采样比直接interpolate效果略好但参数量多如果数据集再小一点换成双线性插值也能跑。最后一层用sigmoid把输出压到 0 到 1配合BCELoss使用。如果换成BCEWithLogitsLoss最后一层就不要加sigmoid否则数值不稳定。2.3 训练循环与损失函数选型DRIVE 数据集有个硬伤血管像素只占全图 10% 左右背景占 90%。直接用BCELoss训练模型会倾向于全预测背景准确率看着有 90%但 Dice 系数接近 0。常见做法是BCE Dice联合损失。def dice_loss(pred, target, smooth1e-6): pred pred.contiguous().view(-1) target target.contiguous().view(-1) intersection (pred * target).sum() return 1 - (2. * intersection smooth) / (pred.sum() target.sum() smooth) def train_one_epoch(model, loader, optimizer, device): model.train() total_loss 0 bce nn.BCELoss() for img, mask, fov in loader: img, mask, fov img.to(device), mask.to(device), fov.to(device) optimizer.zero_grad() pred model(img) # 只在 FOV 有效区域内算损失 pred_fov pred * fov mask_fov mask * fov loss bce(pred_fov, mask_fov) dice_loss(pred_fov, mask_fov) loss.backward() optimizer.step() total_loss loss.item() return total_loss / len(loader)学习率我一般从1e-3开始用Adam优化器跑 50 个 epoch 左右。DRIVE 训练集只有 20 张每个 epoch 迭代很快但要注意过拟合。验证集 Dice 在第 30 个 epoch 之后如果还在涨但训练集 Dice 已经到 0.98那就是过拟合了得加早停或者数据增强。fov掩膜在这里的作用很关键不加的话背景像素会主导梯度模型学不到血管的细结构。3. 评估指标与可视化Dice、IoU 和 ROC 到底看哪个3.1 Dice 系数与 IoU 的计算方式及阈值选择Dice 系数是视网膜血管分割最常用的指标公式是2 * |A ∩ B| / (|A| |B|)。PyTorch 里实现起来很简单但阈值选 0.5 还是 0.4 对结果影响很大。def compute_dice(pred, target, threshold0.5): pred_bin (pred threshold).float() intersection (pred_bin * target).sum() return (2. * intersection) / (pred_bin.sum() target.sum() 1e-6) def compute_iou(pred, target, threshold0.5): pred_bin (pred threshold).float() intersection (pred_bin * target).sum() union pred_bin.sum() target.sum() - intersection return intersection / (union 1e-6)DRIVE 测试集上UNet 不加任何后处理Dice 大概在 0.78 到 0.82 之间。阈值从 0.5 降到 0.4Dice 能涨 0.5 到 1 个点但血管会变粗细血管连成片。我一般会在验证集上扫一遍阈值从 0.3 到 0.7步长 0.05选 Dice 最高的那个。注意测试集不能用这个阈值调否则就是过拟合测试集。3.2 可视化工具把原图、金标准和预测叠在一起看光看指标不够血管分割的很多问题只有可视化才能发现。我习惯写一个visualize.py把原图、手工标注、模型预测、FOV 掩膜拼成一张图。import matplotlib.pyplot as plt def visualize_result(img, mask, pred, fov, save_pathNone): img img.cpu().permute(1, 2, 0).numpy() mask mask.cpu().squeeze().numpy() pred pred.cpu().squeeze().numpy() fov fov.cpu().squeeze().numpy() fig, axes plt.subplots(1, 4, figsize(16, 4)) axes[0].imshow(img) axes[0].set_title(Original) axes[1].imshow(mask, cmapgray) axes[1].set_title(Ground Truth) axes[2].imshow(pred, cmapgray) axes[2].set_title(Prediction) # 叠加显示绿色是金标准红色是预测黄色是重叠 overlay np.zeros((*mask.shape, 3)) overlay[..., 1] mask overlay[..., 0] pred axes[3].imshow(overlay) axes[3].set_title(Overlay (G: GT, R: Pred)) for ax in axes: ax.axis(off) plt.tight_layout() if save_path: plt.savefig(save_path, dpi150) plt.close()叠加图里黄色区域是预测和金标准重叠的部分红色是误分割绿色是漏分割。如果红色集中在血管边缘说明模型对边界不敏感可以加边界损失如果绿色集中在细血管末端说明模型对细小结构欠拟合得加深网络或者加注意力模块。这套可视化工具比盯着 Dice 数字有用得多尤其是调参阶段。4. 避坑与排查DRIVE 训练 UNet 时最容易翻车的 5 个地方4.1 损失不下降Dice 一直卡在 0.1 附近现象训练了 10 个 epochloss 从 0.8 降到 0.7 就下不去了验证集 Dice 在 0.1 到 0.15 之间晃。原因最常见的是标注图没做二值化。DRIVE 的1st_manual是 gif 格式PIL 读出来血管像素值是 1但有些预处理脚本会把它归一化到 0 到 1 之间导致血管像素变成 0.0039 这种极小值模型学到的全是背景。另一个可能是fov掩膜没乘上去背景像素主导了梯度。解决在__getitem__里打印一下mask.max()和mask.min()确认是 1 和 0。如果不是加一行mask (mask 0).astype(np.float32)。同时检查损失函数里有没有乘fov。4.2 验证集 Dice 比训练集低 0.2 以上现象训练集 Dice 0.95验证集只有 0.72差距巨大。原因DRIVE 训练集只有 20 张模型参数量 7M 左右很容易记住训练样本。数据增强只用了翻转多样性不够。解决加随机旋转TF.rotate角度范围 -15 到 15 度、随机亮度对比度扰动TF.adjust_brightness、TF.adjust_contrast、随机裁剪从 512×512 裁到 448×448 再 resize 回去。另外加Dropout2d在编码器最后两层p0.3。早停策略用验证集 Dicepatience 设 10。4.3 显存不够batch size 只能设 1现象RTX 3060 6GB 显存512×512 输入batch size 设 2 就 OOM。原因UNet 在 512×512 分辨率下第一层特征图就是 64×512×512显存占用很大。解决三个方向。一是把features砍到[32, 64, 128, 256]显存降一半Dice 掉 1 到 2 个点。二是用混合精度训练torch.cuda.amp自动把部分计算转成 float16显存省 30% 到 40%。三是把图像裁成 256×256 的 patch 训练推理时再拼回去但拼接处会有缝需要重叠裁剪。4.4 预测结果全是黑色或者全是白色现象模型输出要么全 0 要么全 1Dice 要么 0 要么 1。原因最后一层用了sigmoid但损失函数用了BCEWithLogitsLoss或者反过来。BCEWithLogitsLoss内部自带 sigmoid外面再加一层就重复了输出会被压到 0.5 附近二值化后全是一类。解决检查forward最后一层和损失函数的搭配。用BCELoss就加sigmoid用BCEWithLogitsLoss就不加。我一般统一用BCELoss sigmoid调试的时候直观。4.5 测试集评估时忘了乘 FOV 掩膜现象测试集 Dice 0.85但可视化一看视野外的黑色区域也被算进去了实际血管分割很差。原因DRIVE 测试集的图像有圆形视野视野外是黑色背景。计算 Dice 时如果不乘fov掩膜背景像素会被算成正确预测拉高指标。解决评估函数里统一加pred pred * fov、target target * fov再算 Dice 和 IoU。这个坑很隐蔽因为指标看着不低但实际效果差很多。我一般在测试脚本里强制打印fov.sum()和pred.sum()确认量级对得上。5. 把 UNet 推到 0.85 Dice 以上后处理与注意力模块的实战技巧DRIVE 测试集上原始 UNet 不加任何技巧Dice 大概 0.78 到 0.80。想推到 0.85 以上光调学习率不够得从后处理和网络结构两个方向下手。后处理最有效的是连通域过滤。模型预测出来的二值图里血管应该是连通的但会有一些孤立的噪点。用scipy.ndimage.label找到所有连通域把面积小于 50 像素的去掉Dice 能涨 1 到 1.5 个点。代码很简单from scipy import ndimage def remove_small_objects(pred_bin, min_size50): labeled, num ndimage.label(pred_bin) for i in range(1, num 1): if (labeled i).sum() min_size: pred_bin[labeled i] 0 return pred_bin另一个后处理是形态学闭运算用 3×3 的核把断裂的细血管连起来。cv2.morphologyEx(pred_bin, cv2.MORPH_CLOSE, kernel)迭代 1 到 2 次。注意核不能太大否则血管会粘连。网络结构上加注意力门控是性价比最高的改进。在跳跃连接处加一个AttentionBlock让解码器自动关注血管区域抑制背景。实现上就是在torch.cat之前用解码器特征生成一个注意力权重图乘到编码器特征上。参数量增加不到 5%Dice 能涨 2 到 3 个点。另一个方向是深监督在解码器每一层都接一个 1×1 卷积输出预测和最终输出一起算损失梯度回传更充分对小数据集特别有效。验证方法上我习惯把测试集 20 张图分成 4 组每组 5 张分别算 Dice看方差。如果某组特别低单独把那几张图拿出来可视化大概率是血管特别细或者有病变干扰。这种分组验证比只看平均 Dice 更能发现问题。最后说一个我踩过的坑不要用测试集调阈值和后处理参数。DRIVE 测试集只有 20 张调几次就过拟合了。正确做法是从训练集里切 4 张当验证集所有超参数在验证集上定好测试集只跑一次。这个习惯让我在多个医学图像项目上少走了很多弯路。希望帮到你。本文还有配套的精品资源点击获取

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询