news 2026/9/12 23:39:21

眼底血管分割实战:Unet切片数据集训练与推理全流程解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
眼底血管分割实战:Unet切片数据集训练与推理全流程解析

简介:面向眼底血管分割任务的Unet完整资料包,包含已切片好的数据集、训练代码、推理脚本及训练结果文件。数据集对应眼底血管二分割任务,模型仅训练10个epochs,全局像素准确度达0.95,miou为0.67,若增大训练轮数性能还有提升空间。资源共216个文件,以182个png切片图像为主,另有8个py与14个pyc组成的Python工程、1个pth权重文件、5个xml工程配置及3个txt/readme说明文档,整体压缩包约153.92MB。目前已有269人学习/下载。代码支持多尺度随机缩放训练,utils中的compute_gray函数可自动保存mask灰度值并自适应输出通道,方便扩展多分割项目;训练采用cos学习率衰减,run_results内可查看损失与iou曲线、训练日志及每类指标,推理时只需将图像放入inference目录后运行predict脚本,无需额外参数,适合医学图像分割入门者、需要快速复现Unet分割流程或想训练自定义数据集的研究人员。

1. 眼底血管分割为什么上来就选Unet

给眼底照片做血管分割时,医生最先看的就是血管网络。血管是否变细、迂曲、有没有微动脉瘤,直接对应糖尿病视网膜病变的分级;但血管末端常常只有两三个像素宽,边缘和背景的对比度又低,这类任务必须同时依赖全局语义和局部纹理。Unet几乎是这个任务的默认起点:编码器把感受野逐层扩大,跳跃连接又把浅层边界带回解码器,正好同时保住语义与细节。这也是我拿到一个带切片好的数据集、完整代码和训练结果文件的项目包时,会先去确认数据加载方式和checkpoint落盘策略的原因——这两处决定了换机器后能不能复现原有性能。

2. 切片好的数据集怎么组织:目录约定与 Dataset 读取

拿到项目包,先别急着打开训练脚本。这类项目的目录结构通常逃不出一个套路:train和val下面各有一个img目录和一个mask目录,图片和标签靠文件名一一对应。先把目录结构读明白,后面训练、推理、评估都能少踩一半的坑。这一章就从切片尺寸、同名约定和Dataset实现三个角度把它讲透。

2.1 切片行为什么重要:分辨率、感受野和显存

原始眼底图的分辨率差异很大,从公开数据集的几百像素见方,到竞赛里常见的2000像素以上。如果整张图直接丢给Unet,显存会立刻吃紧,更关键的是感受野问题:Unet做五次下采样后,深层特征图分辨率只有输入图像的1/32,一根两三个像素宽的细血管在深层基本只剩一个响应点,边缘信息全靠跳跃连接从浅层抄回来,所以输入尺寸不能随意缩小,但也不能盲目加大。

因此在拿到“切片好的数据集”时,先确认切片尺寸和切的时候有没有重叠。常见做法是切成512×512或256×256的patch,大图上按固定步长滑窗切出来,再按坐标存回同一命名空间。两种尺寸的取舍见下表:

patch 尺寸显存需求(8G/16G卡可跑batch)细血管保留训练速度
256×25616 / 32一般,末梢血管容易断
512×5124 / 8好,血管连续性明显更稳

我的建议是:只为了快速把代码跑通,256够用;如果目标是要给医生当辅助工具,512起步。切片时的重叠不能省,至少在patch边界的16~32像素要重叠,否则后面推理时会在patch交界处出现假阳性。

2.2 目录约定与同名规则:图片和掩膜如何一一对应

大多数此类项目包的目录结构长这样:

dataset/ ├── train/ │ ├── img/ │ │ ├── 001.png │ │ └── 002.png │ └── mask/ │ ├── 001.png │ └── 002.png └── val/ ├── img/ └── mask/

规则就三条:图片和掩膜严格同名;掩膜是8位灰度PNG,背景0、血管255;不要用JPEG存掩膜,JPEG会把0/255的硬边界压出过渡带。如果包里给的是.npy文件,先确认存的是完整大图还是切片,很多旧脚本会把整张图转成float32的npy,体积比PNG大好几倍,读取时还要额外做一次维度调整。

另一个容易被忽略的文件是fov_mask,也就是视盘区域的掩膜。DRIVE这类公开数据集通常会带一个只覆盖眼底圆形区域的掩膜,评估指标只在掩膜内计算。如果包里正好有这个目录,训练阶段可以不管它,但评估阶段一定要把它当作乘法因子加进去,否则眼底照片四角的黑色区域会把Specificity拉高,Dice也会虚高。

2.3 一个可直接改的torch Dataset:读取切片、归一化和增强

下面这段Dataset是我最常用的一套写法,核心就四点:glob搜集路径、统一读成RGB和灰度、归一化到0~1、在概率上做掩膜二值化,放在单卡训练脚本里能直接跑:

import os from glob import glob import random import numpy as np from PIL import Image import torch from torch.utils.data import Dataset class VesselPatchDataset(Dataset): def __init__(self, img_dir, mask_dir, augment=True): self.img_paths = sorted(glob(os.path.join(img_dir, "*.png"))) self.mask_paths = sorted(glob(os.path.join(mask_dir, "*.png"))) assert len(self.img_paths) == len(self.mask_paths), \ f"img和mask数量不一致: {len(self.img_paths)} vs {len(self.mask_paths)}" self.augment = augment def __len__(self): return len(self.img_paths) def __getitem__(self, idx): img = Image.open(self.img_paths[idx]).convert("RGB") mask = Image.open(self.mask_paths[idx]).convert("L") if self.augment: # 只用90度的整数倍旋转,避免插值破坏单像素细血管 if random.random() > 0.5: img = img.transpose(Image.ROTATE_90) mask = mask.transpose(Image.ROTATE_90) img = np.asarray(img, dtype=np.float32) / 255.0 mask = np.asarray(mask, dtype=np.float32) / 255.0 mask = (mask > 0.5).astype(np.float32) img = torch.from_numpy(img.transpose(2, 0, 1)) mask = torch.from_numpy(mask[None, ...]) return img, mask

代码里img_paths和mask_paths都用sorted(glob(...)),匹配顺序是稳定的;掩膜用convert("L")读成单通道,再除以255并做>0.5的二值化,是为了避免PNG压缩残留的126、254这类脏像素进损失函数。旋转只用90度整数倍而不是随机小角度,因为小角度旋转要插值,对单像素宽的血管是致命的。

这个Dataset默认输入已经是切好的patch,如果拿到的是原始大图,需要再补一个随机裁剪步骤。如果训练时发现iteration跑得飞快但验证指标不动,优先检查mask路径是不是串到了img目录——这是同名目录最容易犯的错。

3. Unet结构要点与训练配置:把完整代码调成能收敛的模型

所谓“完整代码”,通常指三样东西:网络定义、损失函数、训练循环。跑通之前先确认三件事:网络用了几层下采样、损失怎么算、checkpoint按什么频率保存。这一章按这三个顺序讲清楚,顺便把训练环境相关的坑点一起说掉。

3.1 从Unet网络结构图看血管任务的三个关键设计

如果没看过Unet网络结构图,可以先把它理解成一个U形的对称结构:左边编码器每一层做两个3×3卷积加一次池化,通道数从64逐步翻到512;右边解码器先上采样,再和编码器对应层拼接,通道数再降回来。血管分割真正吃到红利的是三个设计。

第一,编码器给了解码器足够的上下文。一根小血管断了,只看局部很难判断该不该连上,编码器把32倍下采样之后的语义带回来,模型才能“猜”出这条血管大概从哪来、到哪去。第二,跳跃连接把第一层和第二层的边缘响应原样拼给解码器,末梢血管的细节就是这样保住的。第三,Unet的参数量对医学图像任务很合适,小数据集上不轻易过拟合,一张8G显存的卡就能训练。

如果你拿到的是带ResNet骨干的变体,注意骨干预训练权重要么有、要么彻底不用;随机初始化一半再冻结一半的训练方式最不稳定,血管分割这种强局部特征任务往往还不如标准Unet好用。

3.2 损失函数选择:BCE、Dice还是两者相加

眼底血管分割的类别不平衡,公开数据集上大约是1:9,也就是说背景像素数量接近血管的十倍。直接上标准BCE时,模型会把几乎所有像素判成背景,因为这样loss已经足够低。所以常见做法是把Dice损失和BCE加在一起:

import torch import torch.nn.functional as F def bce_dice_loss(logits, targets, smooth=1.0): probs = torch.sigmoid(logits) bce = F.binary_cross_entropy_with_logits(logits, targets) inter = (probs * targets).sum(dim=(1, 2, 3)) union = probs.sum(dim=(1, 2, 3)) + targets.sum(dim=(1, 2, 3)) dice = 1 - ((2 * inter + smooth) / (union + smooth)).mean() return 0.5 * bce + 0.5 * dice

bce用binary_cross_entropy_with_logits而不是先过sigmoid再算交叉熵,数值上更稳定;dice的smooth取1.0只在训练早期起作用,不影响最终收敛;0.5和0.5的权重可以按结果微调。我一般会让dice权重大一点,比如0.6 dice、0.4 bce,对小血管更友好。如果预测偏保守,就把dice权重调高;如果预测全是噪声点,把bce权重拉回0.5。

损失特点适合场景
BCE逐像素独立,梯度平稳作为辅助项稳定早期训练
Dice对类别不平衡不敏感,直接优化目标血管这类小目标
BCE+Dice折中,工程上最稳绝大多数血管分割项目

3.3 训练入口:参数表、batch调节和checkpoint保存

把Dataset和损失函数接进训练脚本后,命令行入口我一般长这样:

python train.py \ --img_dir dataset/train/img \ --mask_dir dataset/train/mask \ --val_img_dir dataset/val/img \ --val_mask_dir dataset/val/mask \ --batch_size 8 \ --lr 1e-3 \ --epochs 100 \ --patch_size 512 \ --val_every 5 \ --save_dir checkpoints

这些参数来自常见的训练环境配置:patch为512时,batch 8大约吃16G显存;8G显存的卡就把batch降到4,同时把patch降到320,效果差距不大。学习率用1e-3配Adam跑前20个epoch,是血管分割里最稳的组合之一。

参数推荐值注意
--lr1e-3Adam下偏大,超过50个epoch没降就再降一半
--batch_size4~8显存不够时优先降batch,再降patch
--epochs100大多数公开数据集60~100个epoch足够
--val_every5间隔太久容易错过最佳checkpoint
--patch_size512256能跑但末梢血管会明显变差

checkpoint保存策略坚持两条:始终保留验证集Dice最高的模型,单独存成best_dice.pth;再存一个last.pth用于事故现场排查。你拿到的“训练的结果文件”里如果只有一个模型,先看脚本里保存条件是什么。另外多看一眼train.log里的验证曲线,如果验证Dice在中间epoch就开始下滑,说明要么数据增强太强,要么lr需要在40个epoch附近做一次0.1倍衰减。

提示:换机器复现时,先确认torch版本和CUDA版本跟原log一致,很多加载报错不是代码问题,是版本差异造成的。

4. 加载训练结果文件做推理:切片拼接、阈值和评估

训练跑完,手里有best_dice.pth、last.pth和train.log。这一章讲怎么把这些结果文件变成能看的血管图,再变成能写进文档的数字。最常见的坑是直接把训练脚本里的验证代码拿来做推理,那样往往会做随机增强,还拼了batch,输出顺序对不上原图。

4.1 单张切片推理脚本:从checkpoint加载模型

import torch import numpy as np from PIL import Image model = UNet(in_channels=3, out_channels=1) # 你训练时用的模型类 ckpt = torch.load("checkpoints/best_dice.pth", map_location="cpu") state = ckpt.get("model_state_dict", ckpt) # 去掉DataParallel加的前缀 state = {k.replace("module.", ""): v for k, v in state.items()} model.load_state_dict(state) model.eval() img = Image.open("dataset/val/img/001.png").convert("RGB") x = np.asarray(img, dtype=np.float32) / 255.0 x = torch.from_numpy(x.transpose(2, 0, 1)).unsqueeze(0) with torch.no_grad(): logits = model(x) prob = torch.sigmoid(logits)[0, 0].numpy() Image.fromarray((prob * 255).astype(np.uint8)).save("out_prob.png")

这里有几个关键点。加载前先检查ckpt里的key是不是带module.前缀,用了DataParallel就必须去掉;map_location要跟当前环境一致,半精度权重加载时尤其容易报错;推理前必须调model.eval(),否则带BatchNorm的模型结果会抖动。输出是sigmoid概率而不是阈值图,阈值后处理要放到统计指标时统一做,不要在推理阶段提前二值化。

4.2 滑窗重叠拼接:消除切片边界断裂

训练和推理的区别在拼接这里体现得最明显。训练时可以随机裁剪,因为数据增强了多样性;推理时必须保证每个patch的位置能映射回原图。直接用非重叠滑窗会把血管在patch边界处切成两段,重叠拼接是常见解法:重叠区域被多个patch推理,概率做平均,边界自然平滑。

def sliding_predict(model, img, patch_size=512, overlap=32, device="cuda"): h, w = img.shape[:2] step = patch_size - overlap # 把图pad到能被step整除,避免漏掉右边界和下边界 pad_h = (patch_size - h % step) % step pad_w = (patch_size - w % step) % step img = np.pad(img, ((0, pad_h), (0, pad_w), (0, 0)), mode="reflect") prob_sum = np.zeros_like(img[..., 0], dtype=np.float32) weight = np.zeros_like(img[..., 0], dtype=np.float32) for y in range(0, img.shape[0] - patch_size + 1, step): for x in range(0, img.shape[1] - patch_size + 1, step): crop = img[y:y + patch_size, x:x + patch_size].astype(np.float32) / 255.0 x_t = torch.from_numpy(crop.transpose(2, 0, 1)).unsqueeze(0).to(device) with torch.no_grad(): out = torch.sigmoid(model(x_t))[0, 0].cpu().numpy() prob_sum[y:y + patch_size, x:x + patch_size] += out weight[y:y + patch_size, x:x + patch_size] += 1 prob = prob_sum / np.maximum(weight, 1.0) return prob[:h, :w]

重叠区会累加多次预测,weight矩阵记录每个像素被预测了多少次,最后做除法取平均。overlap设16到32比较合适,太小抑制不了接缝,太大计算量翻倍。如果推理显存紧张,可以把patch_size调小,同时把overlap等比调小,但不要只改patch_size不改overlap。

4.3 评估指标怎么算:Dice、Sensitivity和Specificity

评估时要注意阈值的选择。不要默认0.5,常见做法是在验证集上画出阈值和指标的关系曲线,再取Youden index对应的阈值。四个最常用的指标定义如下:

指标公式关心的问题
Dice2TP / (2TP + FP + FN)血管区域整体重合度
IoUTP / (TP + FP + FN)与Dice类似,分母略大
SensitivityTP / (TP + FN)末梢小血管漏检多少
SpecificityTN / (TN + FP)背景噪声有多少

计算指标的代码可以直接把四个指标合在一起:

def compute_metrics(pred_mask, gt_mask): pred_mask = pred_mask.astype(bool) gt_mask = gt_mask.astype(bool) tp = (pred_mask & gt_mask).sum() tn = (~pred_mask & ~gt_mask).sum() fp = (pred_mask & ~gt_mask).sum() fn = (~pred_mask & gt_mask).sum() eps = 1e-6 return { "dice": 2 * tp / (2 * tp + fp + fn + eps), "iou": tp / (tp + fp + fn + eps), "sens": tp / (tp + fn + eps), "spec": tn / (tn + fp + eps), "acc": (tp + tn) / (tp + tn + fp + fn + eps), }

如果数据集带fov_mask,在调用这个函数之前先对pred_mask和gt_mask做一次& fov操作。没有fov就按全图算,但报告里要注明,因为背景黑色区域会把Specificity顶得很高,不同数据集之间的数字不具备可比性。

5. 让血管不断裂的距离图监督:给Unet加辅助头

如果前面几步都跑通了,最常见的残留问题是末梢血管断成一截截。BCE和Dice都是像素级损失,它们不惩罚结构上的“断点”:只要断裂的像素数量占比低,整体损失看起来依然正常。这时候可以上一个成本很低的Unet模型改进:距离图监督。

先用骨架提取把血管压成中心线,再计算每个骨架像素到血管边缘的距离:

import numpy as np from scipy.ndimage import distance_transform_edt from skimage.morphology import skeletonize def vessel_dist_map(mask): bin_mask = (mask > 0.5).astype(np.uint8) skeleton = skeletonize(bin_mask) dist = distance_transform_edt(bin_mask) return (dist * skeleton).astype(np.float32)

mask是二值标签,skeleton是骨架提取结果;distance_transform_edt算出每个血管像素到背景的最近距离,乘上skeleton之后,只有中心线上的点保留距离值,背景全是0。网络在分割头之外再接一个单通道输出预测这个距离图,相当于给每个血管像素加了一个“你离边缘有多远”的连续目标,监督信号比二值掩膜强得多。

训练时总损失这样拼:

seg_loss = bce_dice_loss(logits, mask) dist_loss = F.smooth_l1_loss(dist_logits, dist_gt) loss = 0.5 * seg_loss + 0.5 * dist_loss

smooth_l1_loss对离群点比MSE稳,适合距离图这种偶尔出现大值的回归目标。分担比例从0.5:0.5起步,如果分割Dice掉得厉害,把dist_weight降到0.3再试。

需要注意骨架提取本身可能把血管末端磨掉,导致末端距离值失真。训练前把距离图可视化一次,看骨架末端是不是缩了一截;skimage的skeletonize对45度走向的血管保留更完整,scipy的morphological_thinning对斜向血管偏保守,两种都可以试。评估时只看分割结果在断裂处有没有变少,不看距离图本身的指标。如果距离图loss降了但分割没改善,先确认辅助头是在同一个decoder上多接了一层卷积,还是从encoder最深层直接引出来——前者通常更稳。

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/12 23:39:06

鸿蒙部署 MicroG 完整教程:解决 Google 服务签名问题

鸿蒙部署 MicroG 完整教程:解决 Google 服务签名问题 【免费下载链接】GmsCore Free implementation of Play Services 项目地址: https://gitcode.com/GitHub_Trending/gm/GmsCore 如果你在一台鸿蒙(HarmonyOS)设备上装过 MicroG&…

作者头像 李华
网站建设 2026/9/12 23:38:39

React Native列表在OpenHarmony上的高性能封装实践

1. 为什么要在OpenHarmony上重新造List这个轮子先说结论:React Native在OpenHarmony上跑通Hello World只是第一步,真正决定能不能上生产的是列表页。FlatList在Android和iOS上表现稳定,但换到OpenHarmony环境后,问题不是“性能差一…

作者头像 李华
网站建设 2026/9/12 23:38:28

CANfestival移植实战:STM32F1上实现CANopen对象字典与PDO/SDO调试

简介:基于CANfestival的CANopen协议在STM32F1系列单片机上的实现,是一份面向嵌入式开发工程师的完整工程资源,解决CANopen协议栈在STM32F1平台下的移植与集成问题。资源共931个文件,压缩包大小28.8MB,包含大量C语言源码…

作者头像 李华
网站建设 2026/9/12 23:37:32

碎纸片拼接:基于TSP建模的组合优化方法

简介:本资源是一项将旅行商问题(TSP)建模思想应用于碎纸片图像拼接复原的MATLAB优化实践项目,面向具备基础图像处理与数学建模能力的本科生、研究生及算法爱好者,解决非结构化纸质文档碎片的自动排序与重建难题。压缩包…

作者头像 李华
网站建设 2026/9/12 23:36:16

10 分钟跑通第一个测试:pytest 入门完整教程

10 分钟跑通第一个测试:pytest 入门完整教程 【免费下载链接】pytest The pytest framework makes it easy to write small tests, yet scales to support complex functional testing 项目地址: https://gitcode.com/GitHub_Trending/py/pytest pytest 是一…

作者头像 李华