☰
基于Unet的眼底血管分割实战:数据集切片、训练与推理全流程
2026/9/28 5:44:58 网站建设 项目流程

简介:本资源面向医学图像分割初学者与深度学习实践者,提供一套基于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 眼底血管分割数据集长什么样,切片为什么绕不开

公开的眼底血管分割数据集,常见的是 DRIVE(40 张 565×584)、CHASE_DB1(28 张 999×960)、STARE(20 张 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, size=64, train=True): 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(p=0.5), A.VerticalFlip(p=0.5), A.RandomRotate90(p=0.5), A.ElasticTransform(alpha=1, sigma=50, p=0.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(image=img, mask=mask) 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', train=True) loader = DataLoader(ds, batch_size=8, shuffle=True, num_workers=2) 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, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) def forward(self, x): return self.net(x) class UNet(nn.Module): def __init__(self, in_ch=3, base=32): 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, stride=2) 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, stride=2) 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, stride=2) self.c2 = nn.Conv2d(base * 2, base * 2, 1) self.d2 = DoubleConv(base * 4, base * 2) self.up1 = nn.ConvTranspose2d(base * 2, base, 2, stride=2) 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)], dim=1) x = self.d4(x) x = self.up3(x) x = torch.cat([x, self.c3(e3)], dim=1) x = self.d3(x) x = self.up2(x) x = torch.cat([x, self.c2(e2)], dim=1) x = self.d2(x) x = self.up1(x) x = torch.cat([x, self.c1(e1)], dim=1) x = self.d1(x) return self.out(x)

逻辑说明:base=32是显存和精度的折中,如果你只有 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 * Dice,BCE 用pos_weight再压一下正样本。优化器用 AdamW,学习率 1e-3,weight 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_ch=3, base=32).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=100) bce = torch.nn.BCEWithLogitsLoss(pos_weight=torch.tensor([5.0]).to(device)) def dice_loss(logits, target, eps=1e-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', train=True) train_loader = DataLoader(train_ds, batch_size=8, shuffle=True, num_workers=2) 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(f'epoch {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()}, f'ckpt_epoch{epoch}.pth')

逻辑说明:pos_weight=5.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, patch=64, stride=32, device='cuda'): 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), dtype=np.float32) weight = np.zeros((h, w), dtype=np.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:y+patch, x:x+patch] 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:y+patch, x:x+patch] += out * gauss weight[y:y+patch, x:x+patch] += gauss prob = prob / np.maximum(weight, 1e-6) return (prob > 0.5).astype(np.uint8) * 255

逻辑说明:stride=32是 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.83
AUC像素级排序能力0.97 到 0.98
Sensitivity血管像素召回率0.78 到 0.82
Specificity背景像素正确率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 左右,想再往上走,我一般会从三个方向动手。第一是深监督,在解码器每一层加一个辅助输出,用同样的 Dice+BCE 监督,让浅层也学到血管结构,这对细血管召回提升最明显,通常能加 1 到 2 个点。第二是注意力跳连,把原始 concat 换成加一个通道注意力,让解码器自己挑编码器里有用的特征,代码改动很小,就是在c4/c3/c2/c1后面接一个 SE block。第三是后处理,对二值掩膜做一次形态学闭运算再细化,能补上断裂的细血管,但核别开大,3×3 就够了,开大了粗血管会粘连。

验证这些改动值不值得做,我的习惯是固定同一个验证集,每次只改一个变量,跑三次取平均,别一次叠三个模块然后说不清是谁的功劳。训练结果文件按exp_name_epoch_dice.pth命名,日志单独存,过两周回头看还能对上号。这套流程我用了很久,最大的教训就是别急着换模型,先把数据、损失、推理这三块抠干净,Unet 在眼底血管分割上的上限比很多人想的高。希望帮到你。

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

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

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

立即咨询