Unet眼底血管分割实战:从数据切片到训练避坑
2026/9/24 20:29:27 网站建设 项目流程

简介:基于U-Net的眼底血管分割完整项目,面向医学图像处理与深度学习入门者,提供从数据集到训练、推理、结果分析的一站式流程。整套7z压缩包内共216个文件,包体约153.92MB,其中182个PNG图像用于训练与预测展示,8个Python脚本及配套pyc文件实现多尺度训练、IoU/损失曲线绘制与一键推理,pth权重文件保存了训练10个epochs的最优模型,txt文本记录类别权重与训练日志,readme可指导环境配置与自定义数据接入。目前已有269人学习下载。项目在10个epochs下即达到全局像素准确率0.95、mIoU 0.67,若加大epoch性能可进一步提升;代码采用cos学习率衰减,自动为Unet输出通道数适配二分割,utils模块会把mask灰度值写入txt,还可在日志中查看各类别IoU、recall、precision。将待推理图片放入inference目录并运行predict脚本即可,无需手动设参,适合快速复现和二次开发。

1. 从一张眼底照片里找血管:Unet 为什么成了这个任务的默认答案

眼底血管分割不是一个新问题,但直到 Unet 出现之前,它都处在一种“能做但不实用”的状态。传统方法靠滤波器和形态学操作提取血管,光照不均、病灶遮挡、毛细血管太细这三座大山压下来,分割结果总是断的。Unet 在 2015 年被提出后,医学图像分割几乎被它统一了——这不是因为它结构多花哨,而是它把“位置信息”和“语义信息”焊在了一起,这在血管这种极度依赖局部细节又需要全局上下文的任务里是致命的优势。

这套方案里包含切片好的数据集、完整代码和训练结果文件,本质上是一个可以直接复现的整链路:从原始眼底图到像素级血管掩膜,Unet 在其中承担的不仅是特征提取,更是一个可训练的空间滤波器。适合谁用?做医学影像分析的初学者、需要快速搭一套分割基线的研究生,以及想在自己的数据集上验证 Unet 效果的算法工程师。本文不绕弯子,直接拆网络结构、数据切片、训练参数和踩坑记录,让你从零把这条链路跑通。

2. Unet 的骨架与眼底血管任务的特殊性:为什么 skip connection 在这里是刚需

2.1 编码器-解码器结构:血管分割需要多大感受野,又需要多细的细节

Unet 的核心架构是一个对称的 U 形:左侧编码器逐层下采样压缩空间分辨率,右侧解码器逐层上采样恢复分辨率。编码器每降一次分辨率,特征图的通道数翻倍,这意味着网络在高层能看到更大的视野——对于眼底血管来说,大血管的走向、视盘附近的血管弧度、病变区域的遮挡关系都需要这种全局信息才能判断。

但光有全局不够。血管分割的难点在于,微细血管的宽度只有 1 到 3 个像素,如果单纯走“编码-解码”的瓶颈结构,细节在逐层池化中早就磨没了。就像你先把一张照片缩小成缩略图,再放大回去,边缘信息是回不来的。下采样是信息损失的过程,而上采样是信息重建的过程,重建质量决定了毛细血管的连续性。

这里就是 Unet 与普通全卷积网络最本质的差别:skip connection。编码器每一层的特征图,不但在本层继续往下走,还会通过横向连接直接拼接到解码器对应层上。这相当于解码器在恢复细节时,直接把原始分辨率下的边缘特征拿过来用,不需要从模糊的高层特征里“猜”。对于血管这种对边界极其敏感的结构,skip connection 不是锦上添花,而是刚需。

2.2 眼底血管的类不均衡问题:血管像素占比只有 10% 左右

我的经验里,第一次跑血管分割最容易翻车的地方不是网络结构,而是损失函数。眼底图像里血管区域通常只占整幅图像的 8% 到 12%,剩下的全是背景。如果你直接拿普通的交叉熵损失去训练,网络会发现“全部预测为背景”就能拿到 90% 以上的准确率,训练过程看起来 loss 在降,实际分割结果是一片黑。

这是语义分割里经典的类不均衡问题,在血管任务上尤其极端。解决办法有两个方向:一是用加权交叉熵,给前景像素更高权重;二是直接用 Dice Loss 或 Focal Loss。我在实际项目中倾向 Dice Loss,因为它直接优化的是分割结果的重叠率,比加权交叉熵更贴近最终目标。但纯 Dice Loss 在小目标上收敛不稳定,常见做法是 Dice Loss 和交叉熵按一定比例叠加,比如0.5 * dice_loss + 0.5 * bce_loss

另一个需要注意的点是数据增强策略。眼底图像有固定的生理结构,视盘通常在图像一侧,血管从视盘向外辐射。做水平翻转和垂直翻转时,语义不会改变,但旋转角度需要谨慎——眼底图像没有绝对的“上方”概念,但旋转超过 90 度会对血管的走向分布产生影响。我一般只用小角度旋转、翻转和弹性形变的组合,不做 90 度整数倍旋转。

2.3 Unet 的经典变体取舍:plain Unet 还是 Unet++

这个项目标题直接写的是“Unet”,没有加++或其他后缀,说明目标是用最基础的结构跑通任务。但你在复现时会面临一个选择:是照搬原始论文的 plain Unet,还是用有密集跳跃连接的 Unet++。

从训练资源的维度看,plain Unet 的参数量在 3100 万左右(取决于编码器深度),Unet++ 因为密集连接会额外增加约 20% 到 30% 的参数量。如果你的 GPU 显存有限,或者数据量本身不大,plain Unet 反而更容易训练。Unet++ 的优势体现在多尺度特征融合上,它在皮肤病变分割和肺结节分割上比 plain Unet 有稳定提升,但在眼底血管上优势并不绝对。

我的建议是:先跑通 plain Unet,把数据、训练、评估这条链路摸清。等你有了基线结果,再考虑改成 Unet++ 或者加注意力模块(比如 Attention Unet)做对比实验。直接上复杂模型而不理解基础结构的边界,后面排查问题会非常痛苦。

3. 切片好的数据集怎么用:从原始眼底图到训练样本的完整转换

3.1 理解切片数据的组织方式:训练集、验证集、标签的对应关系

这个方案的核心资产之一是“切片好的数据集”。所谓切片,在眼底血管分割里通常指两种操作:一是把原始大图切成若干小 patch,因为整张眼底图动辄 500×500 甚至更高分辨率,直接整图训练对显存要求极高;二是按比例划分训练集和验证集,确保评估时用的是模型没见过的数据。

拿到切片后的数据集,先别急着训练。我一般会做一个完整性检查,确认每个训练样本都有对应的标签图,且文件名能对应上。切片数据的目录结构通常是这样的:

dataset/ ├── train/ │ ├── images/ │ │ ├── 01_patch_0.png │ │ ├── 01_patch_1.png │ └── masks/ │ ├── 01_patch_0.png │ └── 01_patch_1.png ├── val/ │ ├── images/ │ └── masks/

先确认标签图的数值范围。眼底血管分割的标签图应该是二值图,只有 0(背景)和 255(血管)两种像素值,但有些数据集会用 0 和 1。训练时要在代码里统一转换,比如mask = mask // 255,否则 loss 计算会出问题。这一步看起来简单,但格式不统一导致的训练崩溃是最常见的前期事故。

3.2 训练集与测试集不能有病人重叠:这是数据切片的铁律

这是切片数据使用时最容易踩的坑。如果原始数据来自多张眼底图,比如 DRIVE 数据集有 40 张图,你不能把所有图像混合后随机切 patch 再划分训练测试——同一个病人的多个 patch 可能同时出现在训练集和测试集里,这会导致模型在测试时“见过”部分病灶信息,评估指标虚高。

正确的做法是:先按图像划分病人集合,再从每个病人的图像里切 patch。例如有 40 张原始图,按 70/30 比例划分训练和测试,即先选 28 张图的 patch 进训练集,剩余 12 张图的 patch 进测试集。切片时训练集和测试集分开切,不共享任何像素。如果你的数据集已经有明确的划分说明,就严格按照说明来。

3.3 自己写切片脚本:patch 大小、重叠率与切片数量的三角关系

如果数据集给的是原始眼底图而不是切片,你需要自己写切片脚本。patch 大小的选择直接影响训练效果和显存占用。patch 越大,模型看到的空间范围越大,上下文信息越足,但显存占用跟随平方级增长;patch 越小,样本数越多(切片数量越多),但可能截断血管的连续性。

以 DRIVE 数据集的 565×584 原始图为例,常见配置是切 256×256 或 128×128 的 patch。256 的 patch 在 8GB 显存上配 batch size 4 基本是上限。下面是一个标准的重叠切片脚本:

import cv2 import numpy as np import os def sliding_window_crop(image, mask, patch_size=256, stride=128): """ 对单张眼底图和对应标签做滑窗裁剪 stride 小于 patch_size 时产生重叠,增加样本量并保留边界连续性 """ h, w = image.shape[:2] patches_img, patches_mask = [], [] for y in range(0, h - patch_size + 1, stride): for x in range(0, w - patch_size + 1, stride): img_patch = image[y:y + patch_size, x:x + patch_size] mask_patch = mask[y:y + patch_size, x:x + patch_size] patches_img.append(img_patch) patches_mask.append(mask_patch) # 如果边缘剩余区域不足 patch_size,做右下角补齐采样 if (h - patch_size) % stride != 0 or (w - patch_size) % stride != 0: for y in [h - patch_size]: for x in [w - patch_size]: if y >= 0 and x >= 0: patches_img.append(image[y:y + patch_size, x:x + patch_size]) patches_mask.append(mask[y:y + patch_size, x:x + patch_size]) return patches_img, patches_mask # 使用示例 image = cv2.imread('raw/01_training.png') # BGR顺序读取 mask = cv2.imread('raw/01_manual1.png', cv2.IMREAD_GRAYSCALE) # 标签以灰度图读取 patches_img, patches_mask = sliding_window_crop(image, mask, patch_size=256, stride=128) os.makedirs('train/images', exist_ok=True) os.makedirs('train/masks', exist_ok=True) for i, (img_p, mask_p) in enumerate(zip(patches_img, patches_mask)): cv2.imwrite(f'train/images/01_patch_{i}.png', img_p) cv2.imwrite(f'train/masks/01_patch_{i}.png', mask_p)

这个脚本的关键参数是stride。stride 等于 patch_size 时不重叠,样本量最小;stride 设为 patch_size 的一半时,样本量约为不重叠的 4 倍,且相邻 patch 有 50% 的重叠区域,这能有效缓解血管在 patch 边界被截断的问题。我的实践是:如果原始图只有 20 到 40 张,必须用重叠切片扩样本;如果原始图超过 100 张,不重叠切片配合数据增强也够用。

4. 完整代码怎么跑通:从 DataLoader 到训练主循环的每个细节

4.1 自定义 Dataset:加载切片图像与掩膜,实时做数据增强

拿到切片数据后,第一步是写一个 PyTorch 的 Dataset 类,负责把图像和标签成对加载。这里要特别注意一个实际问题:眼底图和标签的读取方式不同。原图是 RGB 三通道,用cv2.imread默认读成 BGR,要做颜色通道转换;标签是单通道灰度图,读取时要用cv2.IMREAD_GRAYSCALE,并且把像素值归一化到 0-1 区间。

数据增强应该在 Dataset 里做,而不是提前做离线增强。离线增强会把数据集物理体积膨胀好几倍,而且无法在训练中随机变化。在线增强每次 epoch 产生不同的增强结果,相当于模型看到的是“无限”数据集。学术界的标准做法是原图上做随机翻转、旋转和弹性形变,但标签图做完全相同的变换——必须用同一个随机种子,否则图像和标签错位,训练会直接崩掉。

import torch from torch.utils.data import Dataset import cv2 import numpy as np import albumentations as A class VesselDataset(Dataset): """眼底血管分割数据集,加载切片后的图像与对应的二值标签""" def __init__(self, image_paths, mask_paths, augment=False): self.image_paths = image_paths self.mask_paths = mask_paths self.augment = augment self.transform = A.Compose([ A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5), A.RandomBrightnessContrast(p=0.3), A.ElasticTransform(alpha=1, sigma=50, alpha_affine=50, p=0.2), ]) def __len__(self): return len(self.image_paths) def __getitem__(self, idx): image = cv2.imread(self.image_paths[idx]) image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) mask = cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE) # 标签像素值统一归一到 0/1,二值化防止出现中间值 mask = (mask > 127).astype(np.float32) if self.augment: augmented = self.transform(image=image, mask=mask) image = augmented['image'] mask = augmented['mask'] # HWC -> CHW,像素值归一到 [0, 1] image = image.transpose(2, 0, 1).astype(np.float32) / 255.0 mask = mask[np.newaxis, :, :] return torch.from_numpy(image), torch.from_numpy(mask)

这里用 albumentations 库做增强,它最方便的地方是自动保证 image 和 mask 应用完全相同的随机变换,不用手动同步随机种子。ElasticTransform是医学图像分割里很常用的一种增强,模拟血管的形变,对泛化性有帮助,但注意弹性形变强度不要太大,alpha 超过 2 会让细血管扭曲得过于夸张。

4.2 定义 Unet 模型:用现成实现还是自己搭

实际上,经验丰富的从业者都不会自己从头写 Unet,因为网上已经有大量经过验证的实现。但这个项目标题强调“完整代码”,意味着代码包里的 Unet 结构就是你要用的模型,不必自己重写。你需要做的事情是理解这个模型的输入输出维度以及如何加载预训练权重。

阅读代码包里的模型定义时,重点关注几个地方:输入图像的通道数(眼底图是 3 通道)、第一层卷积的 kernel size、下采样次数(通常 4 次)。确认模型接受 256×256 的输入后,把模型实例化并打印参数量:

import torch from models.unet import UNet # 假设代码包里的模型文件 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = UNet(in_channels=3, out_channels=1, init_features=32).to(device) model = torch.nn.DataParallel(model) # 单机多卡时开启 total_params = sum(p.numel() for p in model.parameters()) trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad) print(f'Total params: {total_params / 1e6:.2f}M, Trainable: {trainable_params / 1e6:.2f}M')

init_features=32表示编码器第一层的初始通道数,每下采样一次翻倍:32 → 64 → 128 → 256 → 512。这个值越大模型容量越大,但小数据集上容易过拟合。如果数据集切片后只有几百张 patch,32 是安全的起点;如果数据量上千,可以尝试 64。

4.3 损失函数与训练主循环:Dice Loss 实现与学习率调度

训练主循环是整个链路的核心。损失函数我推荐 Dice Loss 与二元交叉熵的组合。Dice Loss 的实现有几个细节:需要在 batch 维度和空间维度上同时求交集与并集;分母加上平滑项防止除零;用torch.sigmoid把 logits 转成概率后计算。

import torch import torch.nn as nn import torch.nn.functional as F class DiceBCELoss(nn.Module): """Dice Loss 与 BCE 的组合损失,兼顾像素级精度与区域重叠率""" def __init__(self, smooth=1e-6): super().__init__() self.smooth = smooth self.bce = nn.BCEWithLogitsLoss() def forward(self, logits, targets): probs = torch.sigmoid(logits) # 展平到 batch 维度计算 Dice probs_flat = probs.view(probs.size(0), -1) targets_flat = targets.view(targets.size(0), -1) intersection = (probs_flat * targets_flat).sum(dim=1) union = probs_flat.sum(dim=1) + targets_flat.sum(dim=1) dice = (2.0 * intersection + self.smooth) / (union + self.smooth) dice_loss = 1 - dice.mean() bce_loss = self.bce(logits, targets) return 0.5 * dice_loss + 0.5 * bce_loss

训练主循环里有一个常见错误:总是忘记在验证时把模型切到eval()模式。这会导致 BatchNorm 在推理时继续用当前 batch 的均值方差,验证指标忽高忽低,误以为模型训练不稳定。另外,训练时保存最优模型不只是看验证集 loss,还要同时看 F1 分数——loss 最低的 epoch 不一定分割效果最好,因为在类不均衡任务上,loss 曲线的分辨力不够。

def train_one_epoch(model, dataloader, optimizer, criterion, device): model.train() total_loss = 0.0 for images, masks in dataloader: images, masks = images.to(device), masks.to(device) optimizer.zero_grad() logits = model(images) loss = criterion(logits, masks) loss.backward() optimizer.step() total_loss += loss.item() return total_loss / len(dataloader) def validate(model, dataloader, criterion, device): model.eval() # eval模式:关闭Dropout,BatchNorm用全局统计量 total_loss = 0.0 fp_counter = 0 # 统计全预测为背景的坏模型 with torch.no_grad(): for images, masks in dataloader: images, masks = images.to(device), masks.to(device) logits = model(images) loss = criterion(logits, masks) total_loss += loss.item() preds = (torch.sigmoid(logits) > 0.5).float() if preds.sum() == 0: fp_counter += 1 return total_loss / len(dataloader), fp_counter

上面的验证函数里我加了一个preds.sum() == 0的计数——如果验证集上模型把所有 patch 都预测为背景,说明发生了类不均衡崩溃,每个 epoch 都会产生一个空预测。这比只看 loss 更早发现问题。

4.4 训练参数选择:学习率、batch size 与 epoch 数量怎么定

训练参数是新手最容易用“默认值”带过但实际上影响非常大的环节。我基于多个眼底血管数据集的训练经验,给出以下推荐起点:

参数推荐值说明
优化器Adamβ1=0.9, β2=0.999,适合医学图像分割这类非凸优化
初始学习率1e-4不可大于 1e-3,否则 Dice Loss 早期会震荡
batch size4-8由显存决定,低于 4 时 BatchNorm 不稳定
学习率调度ReduceLROnPlateau验证 loss 连续 5 epoch 不降则 lr 减半
epoch 数80-120配合早停,patience 设为 15
输入尺寸256×256平衡上下文信息与显存占用

优化器选用 Adam 是因为它在医学图像分割任务上几乎不需要调整动量参数就能收敛得很好。但 Adam 有一个已知缺点:后期收敛不够精细。经验做法是用 Adam 训练前 60 个 epoch,然后切换成带 momentum 的 SGD 做微调,学习率降到 1e-5。这个切换能带来 F1 分数 0.01 到 0.02 的提升,属于性价比极高的调参技巧。

学习率调度我有一个血泪教训:不要只看训练 loss 来降学习率,训练 loss 下降不代表验证 loss 也下降。如果训练 loss 持续下降而验证 loss 停在原地不动,是过拟合信号,此时应该降低学习率而不是继续训练。我一般同时监控训练和验证的 loss 曲线,触发条件只认验证 loss。

5. 训练结果文件怎么解:从指标到可视化,以及 4 个避坑记录

5.1 结果文件里有什么:权重、日志、预测图与指标表

训练完成后的结果文件通常包含以下几类内容,你需要知道每个文件怎么用、看什么:

  • 模型权重:best_model.pth是验证集 F1 最优的权重,last_model.pth是最后一个 epoch 的权重。推理时用best_model.pth,续训时用last_model.pth
  • 训练日志:通常是.txt.csv格式,每个 epoch 记录训练 loss、验证 loss、F1、IOU 等指标。用来画 loss 曲线,观察收敛趋势。
  • 预测图:验证集上的分割结果,对比标签图检查视觉效果,这个是最终裁决者。
  • 指标汇总:F1、IOU、AUC、准确率等。看 F1 和 IOU 为主,Acc 和 Dice 为辅。

加载权重做推理的代码框架如下:

model.load_state_dict(torch.load('checkpoints/best_model.pth', map_location=device)) model.eval() # 单张图推理 with torch.no_grad(): input_tensor = preprocess_image('test/01_test.png') # 归一化+转CHW+加batch维度 logits = model(input_tensor.to(device)) probs = torch.sigmoid(logits) pred = (probs > 0.5).float().cpu().numpy().squeeze() # 保存预测结果 cv2.imwrite('results/01_pred.png', (pred * 255).astype(np.uint8))

阈值 0.5 是默认值,但血管分割的最终预测目检时,你会发现阈值可以调。如果模型输出偏保守(血管被低估),把阈值降到 0.4 或多加一个形态学闭运算,会得到更连续的血管;如果背景噪声很多,阈值升到 0.6。阈值的选择属于后处理调优,不改模型不动训练,只影响预测二值化。

5.2 为什么验证集 F1 高但目检效果差:切片边界断裂与阈值问题

这是最迷惑人的一个现象:指标表上 F1 到了 0.80 以上,看起来不错,但把预测图拼回完整原始图像时,血管在 patch 拼接边界处出现明显断裂。原因在于窗口切片时,血管跨越了 patch 边界,模型在 patch 边缘的预测置信度天然偏低——它看到的上下文不完整。

解决这个问题的标准做法是重叠推理:推理时不是从原始图左上角不重叠地切 patch,而是用带 stride 的滑动窗口切 patch,得到多份重叠的预测概率图,然后在重叠区域取均值。这个策略能有效消除拼接缝,代价是推理时间翻倍。class 里推理如下:

def predict_full_image(model, image, patch_size=256, stride=128, device='cuda'): model.eval() h, w = image.shape[:2] prob_map = np.zeros((h, w, 1), dtype=np.float32) count_map = np.zeros((h, w, 1), dtype=np.float32) for y in range(0, h, stride): for x in range(0, w, stride): # 处理边缘越界 y_start = min(y, h - patch_size) x_start = min(x, w - patch_size) patch = image[y_start:y_start+patch_size, x_start:x_start+patch_size] # 预处理并推理 patch_tensor = torch.from_numpy(patch.transpose(2,0,1) / 255.0).unsqueeze(0).float().to(device) with torch.no_grad(): prob = torch.sigmoid(model(patch_tensor)).cpu().numpy()[0, 0] prob_map[y_start:y_start+patch_size, x_start:x_start+patch_size] += prob[..., np.newaxis] count_map[y_start:y_start+patch_size, x_start:x_start+patch_size] += 1 prob_map /= np.maximum(count_map, 1) return prob_map

count_map防止除零:每个像素点可能被多个 patch 覆盖,除以出现次数等于取平均值。

5.3 避坑记录 1:训练集和验证集随机打乱却忘了先按病人划分

现象:训练过程一切正常,验证集 loss 下降得异常快,F1 高达 0.9 以上,但把模型拿到新的眼底图上推理,效果直线下降。

原因:数据划分时没有按病人分组,同一个病人不同位置的 patch 被同时分进了训练集和验证集。模型在验证时偷看了同一个病人的图像分布,指标虚高,泛化能力被高估。

解决:检查数据划分逻辑,以原始图像 ID 为单位划分训练验证,保证验证集中的所有 patch 都来自训练集中没有出现过的原始图像。如果你的切片脚本是对每张原始图独立切片的,那就需要先划分原始图再切 patch。

5.4 避坑记录 2:BatchNorm 在 batch size 为 1 时训练不收敛

现象:显存不足,把 batch size 降到 1 后,训练 loss 震荡剧烈,验证 loss 完全下不去。

原因:BatchNorm 在一个 batch 只有 1 张图时,均值和方差统计量来自单样本,噪声极大,导致模型训练不稳定。眼底血管分割的显存压力主要来自高分辨率输入,不是 batch 大小本身。

解决:把输入 patch 从 256×256 降到 192×192 或 160×160,保住 batch size 为 4 以上;或者去掉 BatchNorm 换成 InstanceNorm。前者更简单,效果也更稳定。

5.5 避坑记录 3:标签图保存为 JPEG 导致血管边界糊成灰色

现象:训练时 loss 能降,但预测结果里血管边缘有一圈灰色晕影,指标 F1 卡在 0.7 左右上不去。

原因:标签图被保存成了 JPEG 格式,压缩过程引入伪影。血管分割的标签是二值的,但 JPEG 压缩会让边界变成过度的灰色值,虽然(mask > 127)能强行二值化,但边界处原本精确的像素位置已经发生了偏移。

解决:检查数据集文件后缀,确保掩膜是 PNG 或 BMP 等无损格式。如果已经买了这份数据集且部分标签是 JPEG,用cv2.imread读入后做 Otsu 二值化,再逐张看效果。

5.6 避坑记录 4:训练到一半显存溢出(OOM)

现象:训练在某个 epoch 中途突然报CUDA out of memory,重启训练后稳定一段时间又复发。

原因:PyTorch 的显存分配不是随 epoch 线性增长的。中途出现 OOM 常见原因是数据加载时的预处理在 CPU 端用了过多内存管理操作,或者某个 batch 中图包含的血管密度高导致中间特征图激活值剧增。另外,验证阶段如果忘了关梯度,验证时也占用同样显存。

解决:训练和验证共用 GPU 时,验证用with torch.no_grad():包裹;如果显存还是不够,开启gradient_accumulation_steps,每个 batch 的梯度累积 2 步再更新一次,等效于增大了 batch size 但显存占用不涨:

scaler.scale(loss).backward() if (step + 1) % 2 == 0: # 每2步更新一次 scaler.step(optimizer) scaler.update() optimizer.zero_grad()

这里用了混合精度训练(torch.cuda.amp),显存占用能减少约 40%。项目代码里如果没启用 AMP,手动加上是一个性价比极高的优化手段。

6. 进阶验证习惯:把预测图叠加到原图上目检,比任何指标都诚实

训练出了一个看起来指标不错的模型,先别急着生成指标表。我有一条坚持了多年的习惯:把预测的血管轮廓叠加到原始眼底图上,逐张看。自动指标有盲区——F1 综合反映整体重叠率,但看不出血管的拓扑连续性。一个模型的 F1 数值不错,可能只是大血管分割得好,细血管全部断掉。叠加目检能在 10 秒内抓住这种问题。

叠加可视化用 OpenCV 很简单:预测图二值化后,找出血管轮廓,在原图上用绿色描边。

def overlay_prediction(original_image, pred_mask, alpha=0.5, color=(0, 255, 0)): """把预测血管区域以半透明绿色叠加到原图上""" overlay = original_image.copy() mask_bool = pred_mask > 0.5 overlay[mask_bool] = (overlay[mask_bool] * alpha + np.array(color) * (1 - alpha)) return overlay # 目检三种典型情况: # 1. 大血管完整、细血管断裂 → 需要降低阈值或加强数据增强 # 2. 血管连续但背景噪点多 → 需要提高阈值或加后处理 # 3. 视盘边界被误判为血管 → 需要检查训练数据的标签正确性

另一个被严重低估的习惯是看每个病例里全部 slice 或 patch 的预测结果拼接图。切片训练最大的代价是丢失全图上下文。如果你训练和推理都发生在 patch 级别,你永远看不到完整血管树的形态。我会在验证阶段对每个验证集病例做整图推理,把所有 patch 预测拼回去,叠加到原始图上连续观察。

用上述方式审查全部验证集结果后,回到指标表里找一个与目检判断一致的指标。如果你觉得血管边缘太平滑,去检查 Dice 是不是被大血管主导了;如果你觉得毛细血管被漏掉不少,去检查阈值是不是要下探。指标数值是裁判,但它只能告诉你分数,不能告诉你哪里踢得差,叠加可视化才能让你看到失分点。

这个项目的地基是 Unet 这个老牌结构,但真正拉开效果差距的从来不是结构本身,而是对数据切片边界、类不均衡、评估方法这三个环节的处理。我做过的分割项目里,凡是效果不如预期的,几乎都能在这三个方向上找到原因。建议你把代码跑通的第一次结果当作基线,然后针对上面三个环节逐一做对比实验,每次只改一个变量。这种控制变量法虽然古老,但在深度学习里依旧是最可靠的调参思路。

希望这些从数据切片到训练排错的细节能帮到你,祝你的眼底血管分割跑出干净、连续、可信的结果。

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

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

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

立即咨询