基于Unet的眼底血管分割实战:数据集切片、训练与推理全流程
发布时间:2026/9/28 5:45:06来源:尧图网络
简介本资源面向医学图像分割初学者与深度学习实践者提供一套基于U-Net的眼底血管二分类分割完整方案解决从数据准备到模型推理的全流程问题。压缩包共216个文件以182张png切片图像、8个py脚本、1个pth权重文件及若干pyc、xml、txt配置与日志为主整体约153.92MB目录结构清晰便于按训练、推理、结果查看等模块检索。资源已包含切片好的数据集、完整代码与训练结果文件仅训练10个epochs即达到全局像素准确率0.95、mIoU 0.67加大训练轮次后性能可进一步提升。代码层面train脚本支持0.5至1.5倍随机缩放的多尺度训练utils中的compute_gray函数可自动保存mask灰度值并定义输出通道数学习率采用cos衰减损失与IoU曲线、训练日志及最优权重均保存在run_results中可查看各类别IoU、recall、precision等指标推理时只需将图像放入inference目录并运行predict脚本即可。目前已有269人学习适合希望快速复现眼底血管分割或迁移到自有数据的小白用户参考。1. 眼底血管分割这件事为什么 Unet 依然是那条最稳的基线眼底血管分割说白了就是把视网膜彩照里那些细如发丝的血管从背景里抠出来做成一张二值掩膜。它直接服务于糖网分级、动静脉交叉压迫分析、高血压视网膜病变筛查这些下游任务血管掩膜的连续性差一点后面的管径测量和分叉点统计就会跟着崩。很多人一上来就想上 Transformer 或者扩散模型但真到落地Unet 依然是那条最稳的基线结构简单、显存友好、在 DRIVE、CHASE_DB1、STARE 这几个公开集上只要预处理和损失函数调对AUC 能稳定压到 0.97 以上。这篇笔记就围绕「基于 Unet 对眼底血管分割」这个方向把切片好的数据集怎么组织、完整代码怎么搭、训练结果文件怎么读、参数怎么调、坑在哪一条线讲透。适合刚接手医学图像分割的算法同学也适合想把眼底血管分割做成一个可复现基线、再往上叠改进模块的工程师。2. 数据集与切片从原始眼底图到能喂进 Unet 的样本2.1 眼底血管分割数据集长什么样切片为什么绕不开公开的眼底血管分割数据集常见的是 DRIVE40 张 565×584、CHASE_DB128 张 999×960、STARE20 张 700×605。这些集有个共同特点原图分辨率不低但标注只有专家勾的第一视角掩膜血管像素占比通常只有 5% 到 12%正负样本极度不平衡。直接把整张图塞进 Unet会遇到两个问题一是显存吃紧batch size 上不去BN 统计不稳二是细血管在多次下采样后直接消失解码器再怎么跳连也补不回来。所以切片patch几乎是标配。常见做法是把原图切成 48×48 或 64×64 的小块训练时按血管像素比例做正负采样推理时再滑窗拼回整图。切片好的数据集一般会给你三个目录train/images、train/masks、test/images外加一个test/masks用于评估。文件名一一对应掩膜是单通道 0/255 的 PNG这点必须先确认否则后面 loss 会算出玄学数值。提示拿到切片数据集第一件事不是写模型而是用脚本统计正负样本比例和每张图的尺寸确认没有混进灰度图或三通道掩膜。2.2 用 Python 把切片数据集读进来并做增强下面这段代码是我一般会先跑的「体检脚本」既读数据又做基础增强顺便把正负比例打出来。它假设数据集是images/和masks/两个平行目录文件名相同。import os import cv2 import numpy as np import albumentations as A from torch.utils.data import Dataset, DataLoader class VesselDataset(Dataset): def __init__(self, img_dir, mask_dir, size64, trainTrue): self.img_dir img_dir self.mask_dir mask_dir self.names sorted(os.listdir(img_dir)) self.size size self.train train # 训练增强翻转、旋转、弹性形变血管对弹性形变很敏感 self.aug A.Compose([ A.HorizontalFlip(p0.5), A.VerticalFlip(p0.5), A.RandomRotate90(p0.5), A.ElasticTransform(alpha1, sigma50, p0.3), ]) def __len__(self): return len(self.names) def __getitem__(self, idx): name self.names[idx] img cv2.imread(os.path.join(self.img_dir, name), cv2.IMREAD_COLOR) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) mask cv2.imread(os.path.join(self.mask_dir, name), cv2.IMREAD_GRAYSCALE) # 掩膜二值化防止有人存成 0/1 或 0/255 混用 mask (mask 127).astype(np.float32) if self.train: out self.aug(imageimg, maskmask) img, mask out[image], out[mask] img img.astype(np.float32) / 255.0 # 标准化用眼底图常见均值方差别用 ImageNet 的 img (img - np.array([0.485, 0.456, 0.406])) / np.array([0.229, 0.224, 0.225]) img np.transpose(img, (2, 0, 1)) mask np.expand_dims(mask, 0) return img.astype(np.float32), mask.astype(np.float32) if __name__ __main__: ds VesselDataset(dataset/train/images, dataset/train/masks, trainTrue) loader DataLoader(ds, batch_size8, shuffleTrue, num_workers2) imgs, masks next(iter(loader)) print(batch shape:, imgs.shape, masks.shape) print(正样本比例:, masks.mean().item())逻辑说明VesselDataset把图像和掩膜同步读入增强用 albumentations 保证几何变换一致掩膜强制二值化是为了避免标注里出现 128 这种中间值导致 BCE 梯度异常。参数上size要和切片时保持一致ElasticTransform的alpha别开太大超过 2 会把血管拉断反而教坏模型。标准化这里用的是眼底图常用的一组均值方差如果你换数据集建议自己统计一遍别直接套 ImageNet 的。2.3 切片策略与正负采样比例怎么定切片不是随便切。我一般会保证每个 batch 里血管像素占比在 30% 到 50% 之间做法是维护两个索引池血管像素超过阈值的 patch 进正池低于阈值的进负池每个 epoch 按 1:1 或 1:2 采样。这样比纯随机切片的收敛快很多Dice 能早两三个 epoch 起来。切片步长建议取 patch 的一半重叠部分在推理时用高斯权重融合能明显减少拼接缝。注意如果你的切片数据集已经切好先看它的命名规则里有没有带坐标带坐标的可以直接拼回整图不带坐标的推理时只能重新滑窗别硬拼。3. Unet 网络搭建与训练完整代码怎么落地3.1 Unet 结构里真正影响血管分割的三个位置标准 Unet 是编码器四次下采样、解码器四次上采样、四次跳连。放到眼底血管分割上有三个位置决定成败。第一是编码器第一层输入通道 3输出通道别一上来就 64血管太细浅层特征要够密我一般用 32 起步再翻倍。第二是跳连原始 Unet 是直接 concat血管分割里更推荐在 concat 前加一个 1×1 卷积压通道减少解码器负担。第三是输出层二分类用 1 通道加 sigmoid别用 2 通道 softmax后者在极端不平衡下更容易全预测背景。下面是一个可直接跑的 Unet 实现带跳连处的通道压缩。import torch 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.net nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.net(x) class UNet(nn.Module): def __init__(self, in_ch3, base32): super().__init__() # 编码器 self.e1 DoubleConv(in_ch, base) self.e2 DoubleConv(base, base * 2) self.e3 DoubleConv(base * 2, base * 4) self.e4 DoubleConv(base * 4, base * 8) self.pool nn.MaxPool2d(2) # 瓶颈 self.bottleneck DoubleConv(base * 8, base * 16) # 解码器跳连前用 1x1 压通道 self.up4 nn.ConvTranspose2d(base * 16, base * 8, 2, stride2) self.c4 nn.Conv2d(base * 8, base * 8, 1) self.d4 DoubleConv(base * 16, base * 8) self.up3 nn.ConvTranspose2d(base * 8, base * 4, 2, stride2) self.c3 nn.Conv2d(base * 4, base * 4, 1) self.d3 DoubleConv(base * 8, base * 4) self.up2 nn.ConvTranspose2d(base * 4, base * 2, 2, stride2) self.c2 nn.Conv2d(base * 2, base * 2, 1) self.d2 DoubleConv(base * 4, base * 2) self.up1 nn.ConvTranspose2d(base * 2, base, 2, stride2) self.c1 nn.Conv2d(base, base, 1) self.d1 DoubleConv(base * 2, base) self.out nn.Conv2d(base, 1, 1) def forward(self, x): e1 self.e1(x) e2 self.e2(self.pool(e1)) e3 self.e3(self.pool(e2)) e4 self.e4(self.pool(e3)) b self.bottleneck(self.pool(e4)) x self.up4(b) x torch.cat([x, self.c4(e4)], dim1) x self.d4(x) x self.up3(x) x torch.cat([x, self.c3(e3)], dim1) x self.d3(x) x self.up2(x) x torch.cat([x, self.c2(e2)], dim1) x self.d2(x) x self.up1(x) x torch.cat([x, self.c1(e1)], dim1) x self.d1(x) return self.out(x)逻辑说明base32是显存和精度的折中如果你只有 8G 显存切片 64×64、batch 8 完全跑得动。跳连处的c4/c3/c2/c1是 1×1 卷积作用是把编码器特征压到和解码器同通道数再 concat避免解码器第一层卷积核过大。输出层不加 sigmoid因为后面用带 logits 的损失函数更稳。3.2 损失函数与优化器Dice BCE 的组合怎么配血管分割里纯 BCE 会被背景淹没纯 Dice 在早期梯度不稳。我一般用0.5 * BCE 0.5 * DiceBCE 用pos_weight再压一下正样本。优化器用 AdamW学习率 1e-3weight decay 1e-4配合 CosineAnnealing 到 1e-5。下面这段训练循环可以直接抄。import torch from torch.utils.data import DataLoader from dataset import VesselDataset from unet import UNet device torch.device(cuda if torch.cuda.is_available() else cpu) model UNet(in_ch3, base32).to(device) optimizer torch.optim.AdamW(model.parameters(), lr1e-3, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max100) bce torch.nn.BCEWithLogitsLoss(pos_weighttorch.tensor([5.0]).to(device)) def dice_loss(logits, target, eps1e-6): prob torch.sigmoid(logits) inter (prob * target).sum(dim(2, 3)) union prob.sum(dim(2, 3)) target.sum(dim(2, 3)) return 1 - (2 * inter eps) / (union eps) train_ds VesselDataset(dataset/train/images, dataset/train/masks, trainTrue) train_loader DataLoader(train_ds, batch_size8, shuffleTrue, num_workers2) for epoch in range(100): model.train() total_loss 0 for img, mask in train_loader: img, mask img.to(device), mask.to(device) optimizer.zero_grad() logits model(img) loss 0.5 * bce(logits, mask) 0.5 * dice_loss(logits, mask).mean() loss.backward() optimizer.step() total_loss loss.item() scheduler.step() print(fepoch {epoch}, loss {total_loss / len(train_loader):.4f}) # 每 10 个 epoch 存一次结果文件 if epoch % 10 0: torch.save({epoch: epoch, model: model.state_dict(), optimizer: optimizer.state_dict()}, fckpt_epoch{epoch}.pth)逻辑说明pos_weight5.0是根据正样本占比约 10% 反推的如果你的数据集血管更稀疏可以调到 8 到 10。Dice loss 在 batch 维度上先算再平均别在像素维度上直接平均否则小 batch 下方差很大。保存的ckpt里带 epoch、模型和优化器状态这就是标题里说的「训练的结果文件」后面恢复训练或做推理都靠它。3.3 训练结果文件里到底存了什么怎么读训练结果文件一般有两种.pth权重文件和.log日志文件。.pth里我习惯存三样东西model.state_dict()、optimizer.state_dict()、当前 epoch 和最佳指标。读的时候别直接torch.load完就model.load_state_dict先看 key 对不对尤其是你改过网络结构之后key 不匹配会直接报错。日志文件建议每 epoch 追加一行epoch, loss, dice, auc方便后面画曲线定位过拟合点。提示如果结果文件里只有state_dict没有epoch恢复训练时学习率调度会从头开始CosineAnnealing 会直接跳回高学习率这是很多人翻车的地方。4. 推理、评估与踩坑把 Dice 从 0.78 拉到 0.82 的细节4.1 滑窗推理与整图拼接的正确姿势切片训练完推理时要把 patch 拼回整图。常见做法是滑窗步长取 patch 的一半每个像素用高斯权重累加最后除以权重和。这样拼接缝几乎看不见。下面这段是推理脚本的核心部分。import cv2 import numpy as np import torch def infer_full_image(model, img_path, patch64, stride32, devicecuda): model.eval() img cv2.imread(img_path) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0 h, w, _ img.shape prob np.zeros((h, w), dtypenp.float32) weight np.zeros((h, w), dtypenp.float32) # 高斯权重中心高边缘低 g cv2.getGaussianKernel(patch, patch / 4) gauss g g.T for y in range(0, h - patch 1, stride): for x in range(0, w - patch 1, stride): patch_img img[y:ypatch, x:xpatch] patch_img (patch_img - np.array([0.485, 0.456, 0.406])) / np.array([0.229, 0.224, 0.225]) tensor torch.from_numpy(patch_img.transpose(2, 0, 1)).unsqueeze(0).float().to(device) with torch.no_grad(): out torch.sigmoid(model(tensor)).cpu().numpy()[0, 0] prob[y:ypatch, x:xpatch] out * gauss weight[y:ypatch, x:xpatch] gauss prob prob / np.maximum(weight, 1e-6) return (prob 0.5).astype(np.uint8) * 255逻辑说明stride32是 patch 的一半重叠区域靠高斯权重融合。getGaussianKernel的 sigma 取 patch/4太大边缘权重压不下去太小中心过尖。阈值 0.5 是默认值实际调的时候可以在验证集上扫 0.3 到 0.7血管分割往往 0.4 左右能多捞回一些细血管。4.2 评估指标Dice、AUC、敏感度一个都不能少只看 Dice 会被背景骗。眼底血管分割里细血管的敏感度Sensitivity和 AUC 更能反映模型对细小结构的捕捉能力。我一般三个指标一起看Dice 看整体重叠AUC 看排序能力Sensitivity 看细血管召回。如果 Dice 高但 Sensitivity 低说明模型在偷懒只预测粗血管这时候要回去查正样本采样比例和 pos_weight。指标含义血管分割里的合理区间Dice预测与标注重叠度0.80 到 0.83AUC像素级排序能力0.97 到 0.98Sensitivity血管像素召回率0.78 到 0.82Specificity背景像素正确率0.98 以上4.3 避坑与排查五个真实踩过的坑现象一训练 loss 一直降但 Dice 卡在 0.6 不动。原因多半是掩膜没二值化或者图像和掩膜文件名没对齐模型在学噪声。解决跑一遍体检脚本打印每对图像的掩膜唯一值确认只有 0 和 1。现象二验证集 Dice 比训练集低 0.1 以上。这是过拟合切片数据集样本量小的时候特别常见。解决加 ElasticTransform 和随机亮度对比度扰动把 weight decay 提到 1e-3或者直接减小编码器通道数。现象三推理整图出现明显网格缝。滑窗步长等于 patch 大小没有重叠。解决stride 改成 patch 的一半并加高斯权重融合。现象四恢复训练后学习率突然跳高loss 炸一下。结果文件里没存 scheduler 状态。解决保存时把scheduler.state_dict()一起存恢复时先 load 再 step。现象五换数据集后 AUC 掉到 0.9 以下。新数据集的成像设备不同均值和方差变了。解决在新数据集上重新统计均值和方差别硬套旧参数。5. 进阶技巧把 Unet 的细血管召回再往上提一档如果你已经把基线跑到 Dice 0.82 左右想再往上走我一般会从三个方向动手。第一是深监督在解码器每一层加一个辅助输出用同样的 DiceBCE 监督让浅层也学到血管结构这对细血管召回提升最明显通常能加 1 到 2 个点。第二是注意力跳连把原始 concat 换成加一个通道注意力让解码器自己挑编码器里有用的特征代码改动很小就是在c4/c3/c2/c1后面接一个 SE block。第三是后处理对二值掩膜做一次形态学闭运算再细化能补上断裂的细血管但核别开大3×3 就够了开大了粗血管会粘连。验证这些改动值不值得做我的习惯是固定同一个验证集每次只改一个变量跑三次取平均别一次叠三个模块然后说不清是谁的功劳。训练结果文件按exp_name_epoch_dice.pth命名日志单独存过两周回头看还能对上号。这套流程我用了很久最大的教训就是别急着换模型先把数据、损失、推理这三块抠干净Unet 在眼底血管分割上的上限比很多人想的高。希望帮到你。本文还有配套的精品资源点击获取
网站建设高端定制企业官网