U-Net 图像分割实战:基于 deep-learning-for-image-processing 仓库的 DRIVE 视网膜血管分割与 PyTorch 训练部署指南

发布时间:2026/9/30 13:21:38
U-Net 图像分割实战:基于 deep-learning-for-image-processing 仓库的 DRIVE 视网膜血管分割与 PyTorch 训练部署指南
示例工程【免费下载链接】deep-learning-for-image-processingdeep learning for image processing including classification and object-detection etc.项目地址https://gitcode.com/gh_mirrors/de/deep-learning-for-image-processing点击查看免费下载U-Net 是生物医学图像分割领域的经典全卷积网络其对称的编码器-解码器结构与跳跃连接设计能够在有限标注数据下同时保留全局语义与像素级细节。本文以deep-learning-for-image-processing仓库中 pytorch_segmentation/unet 模块为实践主体完整讲解该实现的环境配置、文件结构、DRIVE 视网膜血管分割数据集的准备与均值/方差统计、单机单卡与多卡训练含混合精度与断点续训、Dice 系数评估以及基于训练权重的前景掩码推理流程。读者读完本文后将能独立完成从数据预处理到 U-Net 训练、评估、预测的完整实战闭环并掌握仓库中src/unet.py、train.py、train_multi_GPU.py、my_dataset.py、predict.py等关键文件的底层实现原理。一、项目概述与环境配置本模块以经典的 U-Net 架构为基础主要参考了 milesial/Pytorch-UNet 与 torchvision 等开源实现将其适配到 DRIVE视网膜血管分割数据集上形成了包含数据读取、训练、多 GPU 分布式训练、评估与推理的完整工程。1.1 环境要求根据 README.md运行本项目建议满足以下环境Python3.6 / 3.7 / 3.8PyTorch1.10 及以上训练脚本中大量使用torch.cuda.amp混合精度接口请确保所用版本支持操作系统Ubuntu 或 CentOSWindows 暂不支持多 GPU 训练单卡/CPU 训练可用硬件最好使用 GPU 训练多卡训练通过torchrunPyTorch 1.10 起的官方分布式启动器拉起多进程。具体的依赖版本以 requirements.txt 为准核心依赖如下numpy1.22.0 torch1.13.1 torchvision0.11.1 Pillow说明requirements.txt中torch1.13.1与torchvision0.11.1存在版本组合差异实际安装时建议根据你的 CUDA 版本从 PyTorch 官方渠道安装相互匹配的 torch / torchvision 组合避免二进制不兼容。训练脚本本身对版本的敏感点主要在于torchvision.transforms.InterpolationMode详见下文 transforms 分析。1.2 文件结构pytorch_segmentation/unet/ ├── src/ # 搭建 U-Net 模型的代码unet.py 等 ├── train_utils/ # 训练、验证以及多 GPU 训练相关模块 ├── my_dataset.py # 自定义 Dataset用于读取 DRIVE 数据集 ├── train.py # 以单 GPU 为例的训练脚本 ├── train_multi_GPU.py # 针对使用多 GPU 用户的训练脚本 ├── predict.py # 简易预测脚本使用训练好的权重进行推理 ├── compute_mean_std.py # 统计数据集各通道的均值和标准差 ├── transforms.py # 图像与标签同步变换的数据增强 └── requirements.txt # 依赖清单二、U-Net 网络结构源码剖析2.1 经典 U-Net 架构回顾如上图所示U-Net 整体呈对称的 U 形结构由三部分构成编码器收缩路径每层由「2 个 3×3 卷积 BN ReLU」组成图中蓝色模块Conv 3×3, BN, ReLU层间通过 2×2 最大池化MaxPool 2×2步长 2实现空间尺寸减半、通道数翻倍逐步提取高维语义特征解码器扩展路径先通过上采样恢复空间分辨率再执行「2 个 3×3 卷积 BN ReLU」图中橙色模块修正特征逐层恢复细节跳跃连接灰色虚线将编码器每一层的特征图与解码器对称层级的特征图在通道维度拼接弥补下采样丢失的空间细节——编码器特征保留了像素级位置信息解码器特征具备全局语义融合后兼顾语义准确性与细节精度输出头最后通过 1×1 卷积Conv 1×1将特征映射为类别数通道得到逐像素的分割预测。本仓库实现默认使用双线性插值bilinear interpolation作为上采样方式这一点在 README.md 中明确说明也与上图橙色箭头Bilinear Interpolate对应。2.2 源码实现细节本模块的模型定义位于 src/unet.py通过DoubleConv、Down、Up、OutConv四个基础模块拼装出完整网络DoubleConvsrc/unet.py连续两个Conv2d(3×3, padding1, biasFalse) BatchNorm2d ReLU(inplaceTrue)的序列是编码器/解码器中的基本特征变换单元Downsrc/unet.pyMaxPool2d(2, stride2)后接DoubleConv对应一次下采样Upsrc/unet.py当bilinearTrue时使用nn.Upsample(scale_factor2, modebilinear, align_cornersTrue)通道减半后经DoubleConv当bilinearFalse时改用nn.ConvTranspose2d(in_channels, in_channels//2, kernel_size2, stride2)转置卷积上采样。Up.forward中还对特征图做了一次关键的尺寸对齐处理diff_y x2.size()[2] - x1.size()[2] diff_x x2.size()[3] - x1.size()[3] # padding_left, padding_right, padding_top, padding_bottom x1 F.pad(x1, [diff_x // 2, diff_x - diff_x // 2, diff_y // 2, diff_y - diff_y // 2]) x torch.cat([x2, x1], dim1)当输入宽高不是 2 的整数次幂时编码器下采样会产生尺寸取整误差该段代码通过中心式 padding 将上采样结果补齐到与跳跃连接特征一致保证torch.cat拼接合法OutConvsrc/unet.py单个 1×1 卷积将通道数映射为num_classesUNet主类src/unet.py构造函数参数包括in_channels默认 1、num_classes默认 2、bilinear默认 True与base_c默认 64即第一层基础通道数。前向传播依次执行四次下采样down1~down4与四次上采样up1~up4其中factor 2 if bilinear else 1用于调整瓶颈层与上采样层的通道数双线性上采样不产生可学习参数因此瓶颈通道减半以平衡参数与显存最终返回{out: logits}字典形式输出。在训练脚本中train.py、train_multi_GPU.py模型实例化为UNet(in_channels3, num_classesnum_classes, base_c32)即输入三通道 RGB、基础通道数为 32显存占用更小适合 DRIVE 这种小分辨率任务。值得一提的还有 src/vgg_unet.py 与 src/mobilenet_unet.py 两个变体文件分别以 VGG、MobileNet 作为编码器主干替换默认的双卷积单元为需要更换 backbone 的读者提供了扩展方向本文默认介绍的是标准 U-Net 实现。三、DRIVE 数据集准备与归一化统计3.1 数据集下载与目录结构DRIVEDigital Retinal Images for Vessel Extraction是视网膜血管分割的公开基准数据集。README 提供了两个获取渠道官网https://drive.grand-challenge.org/百度云链接密码 8no8。下载解压后必须保证根目录下存在DRIVE文件夹且内部结构符合 my_dataset.py 的读取约定DRIVE/ ├── training/ │ ├── images/ # 训练图像.tif │ ├── 1st_manual/ # 人工标注血管掩码_manual1.gif │ └── mask/ # ROI 掩码_training_mask.gif └── test/ ├── images/ # 测试图像.tif ├── 1st_manual/ # 人工标注_manual1.gif └── mask/ # ROI 掩码_test_mask.gif3.2 数据集读取逻辑my_dataset.py 中的DriveDataset实现了自定义Dataset根据trainTrue/False拼接DRIVE/training或DRIVE/test路径并逐一校验图像、1st_manual标注与maskROI 文件是否存在缺失即抛出FileNotFoundError__getitem__中读取 RGB 图像将人工标注除以 255 归一化到 0/1并将 ROI 掩码取反255 - mask后与标注相加、裁剪到[0, 255]这样标签中0为背景、255为 ROI 外区域作为ignore_index处理——ROI 之外既不是背景也不是血管应在损失函数中被忽略collate_fnmy_dataset.py由于随机裁剪前图像尺寸不固定批量拼接时按 batch 内最大宽高pad图像填充 0、标签填充 255保证一个 batch 内张量形状一致标签中的 255 恰好与上述 ROI 忽略语义一致。3.3 均值与标准差统计训练脚本中默认使用的归一化参数是mean (0.709, 0.381, 0.224) std (0.127, 0.079, 0.043)这两个三元组并非 ImageNet 的通用值而是由 compute_mean_std.py 在 DRIVE 训练集上统计得到。该脚本遍历DRIVE/training/images下的全部.tif对每张图像只取 ROI 掩码内roi_img 255的像素计算逐通道均值和标准差再对所有图像取平均。这一细节很有价值统计时排除 ROI 之外的黑色边框区域可以避免无关像素拉偏归一化参数这正是 README 中提示「使用 compute_mean_std.py」的目的所在。读者若更换数据集可运行python compute_mean_std.py重新统计后替换训练/预测脚本中的mean、std。四、数据增强与训练变换train.py 与 train_multi_GPU.py 中定义了两套变换训练变换SegmentationPresetTrain依次为RandomResize(min_sizeint(0.5*base_size), max_sizeint(1.2*base_size))、概率各为 0.5 的随机水平翻转与垂直翻转、RandomCrop(crop_size)、ToTensor、Normalize。默认base_size565、crop_size480见 train.py即先把图像随机缩放到边长 282~678 之间再中心裁剪/随机裁剪到 480×480验证变换SegmentationPresetEval仅做ToTensor Normalize不做随机增强。这些变换实现在 transforms.py 中其关键点在于图像与标签必须同步变换Compose将(image, target)二元组依次传入每个变换RandomResize对标签缩放时使用最近邻插值T.InterpolationMode.NEAREST注释中特别提示该枚举在 torchvision 0.9.0 之后才可用旧版本需改用PIL.Image.NEAREST避免双线性插值产生介于 0/1/255 之间的“灰色”伪标签水平/垂直翻转同样对 image 与 target 成对执行。任何针对图像单独做的增强如仅对图像生效的色彩抖动都可能破坏标签对齐这是分割训练中最容易踩的坑。五、单 GPU 训练实战5.1 启动命令与参数确保数据集就绪后单卡/CPU 训练直接运行python train.py --data-path DRIVE根目录其中--data-path必须指向DRIVE 文件夹所在的根目录例如--data-path /path/to/data脚本内部会再拼接DRIVE/training这是 README「注意事项」中反复强调的一点。train.py 通过 argparse 暴露的全部可调参数如下参数默认值说明--data-path./DRIVE 数据集根目录--num-classes1前景类别数不含背景脚本内部自动 1 得到总类别数见 train.py--devicecuda训练设备无 GPU 时自动回退 cpu-b/--batch-size4batch size验证集固定为 1--epochs200总训练轮数--lr0.01初始学习率--momentum0.9SGD 动量--wd/--weight-decay1e-4权重衰减--print-freq1打印日志频率步数--resume空断点续训权重路径--start-epoch0起始轮数--save-bestTrue仅保存 Dice 系数最高的权重--ampFalse是否启用torch.cuda.amp混合精度训练5.2 训练主流程main 函数的执行逻辑可以概括为确定设备cuda不可用则回退cpunum_classes args.num_classes 1背景 前景DRIVE 场景下总类别为 2构建训练/验证DriveDataset与 DataLoader其中num_workers取min(os.cpu_count(), batch_size, 8)创建UNet(in_channels3, num_classesnum_classes, base_c32)使用 SGD 优化器momentum0.9、weight_decay1e-4创建混合精度GradScaler仅当--amp开启时创建按 step 更新的学习率调度器create_lr_scheduler非按 epoch若指定--resume加载 checkpoint 中的模型、优化器、调度器状态与start_epochAMP 时还会恢复 scaler逐 epoch 调用train_one_epoch与evaluate将 loss、lr、Dice 系数写入resultsYYYYMMDD-HHMMSS.txt按--save-best保存save_weights/best_model.pthDice 最高或save_weights/model_{epoch}.pth每轮都存。5.3 损失函数与评估指标训练与评估的核心实现位于 train_utils/train_and_eval.pycriteriontrain_and_eval.py采用CrossEntropyLoss DiceLoss的组合损失。交叉熵通过ignore_index255忽略 ROI 外像素当num_classes 2时背景/前景二分类还会传入loss_weight[1.0, 2.0]见 train_and_eval.py加大前景血管在损失中的权重缓解正负样本极度不均衡的问题DiceLossdice_coefficient_loss.py对 softmax 后的概率图与 one-hot 标签计算 Dice 系数并取1 - dice作为损失其中build_target将忽略像素在 one-hot 后仍标记为ignore_indexdice_coeff对忽略像素进行掩码剔除dice_coefficient_loss.pyevaluatetrain_and_eval.py以ignore_index255构建混淆矩阵ConfusionMatrix与DiceCoefficient在torch.no_grad()下逐 batch 更新指标最后打印混淆矩阵与 Dice 系数create_lr_schedulertrain_and_eval.py实现 warmup poly 式学习率衰减。训练开始时倍率因子从warmup_factor1e-3线性升至 11 个 warmup epoch之后按(1 - progress)^0.9多项式衰减参考 deeplab_v2 的 learning rate policy。注意脚本注释提醒PyTorch 在训练开始前会提前调用一次lr_scheduler.step()编写自定义调度器时需留意这一点。六、多 GPU 分布式训练6.1 启动方式多卡训练使用 PyTorch 官方的torchrun启动器对应 README.md 中的说明torchrun --nproc_per_node8 train_multi_GPU.py--nproc_per_node为使用的 GPU 数量。如需指定具体 GPU 设备可在指令前添加CUDA_VISIBLE_DEVICES例如只用物理设备中的第 1 块和第 4 块CUDA_VISIBLE_DEVICES0,3 torchrun --nproc_per_node2 train_multi_GPU.py6.2 与单卡训练脚本的差异train_multi_GPU.py 与train.py共享大部分逻辑差异集中在分布式相关部分通过init_distributed_mode(args)train_utils/distributed_utils.py初始化进程组支持env://等--dist-url方式使用DistributedSampler切分数据每个 epoch 前需调用train_sampler.set_epoch(epoch)打乱顺序见 train_multi_GPU.py模型经torch.nn.parallel.DistributedDataParallel包装--sync-bn参数可开启SyncBatchNorm多卡间同步 BN会降低训练速度权重保存与日志写入只在主进程args.rank in [-1, 0]执行通过save_on_master落盘额外提供--test-only仅测试不训练、--output-dir默认./multi_train、--workers、--world-size等参数多卡场景下--lr建议按 GPU 数量同比放大脚本注释使用 n 块 GPU 建议学习率乘以 nREADME 提醒Windows 暂不支持多 GPU 训练分布式训练请使用 Linux。七、模型预测与推理训练完成后可使用 predict.py 对测试图像进行推理。脚本核心流程如下配置输入将weights_path设置为你训练生成的权重路径默认./save_weights/best_model.pthimg_path指向测试图像默认./DRIVE/test/images/01_test.tifroi_mask_path指向对应的 ROI 掩码加载模型UNet(in_channels3, num_classesclasses1, base_c32)其中classes 1不含背景并从 checkpoint 中取出model键加载权重checkpoint 还包含优化器、调度器等训练状态预处理使用与训练一致的mean/std做ToTensor Normalize并增加 batch 维度推理model.eval()后先向网络送入一张与输入等尺寸的零张量init_img完成一次前向再对真实输入计时推理这一步常见于显存按需分配的 GPU 环境可避免首次前向时的 CUDA 内核初始化时间计入推理耗时后处理取output[out].argmax(1)得到预测类别图将前景类别 1像素置 255白色并将 ROI 掩码之外的像素置 0黑色最后保存为test_result.png。该脚本同时打印单张图像的inference time可用于粗略评估模型在目标硬件上的推理延迟。八、训练过程记录与结果追踪两种训练脚本都会在运行时生成形如results20220109-165837.txt的记录文件仓库中已有示例 results20220109-165837.txt。每个 epoch 会追加记录train_loss该 epoch 的平均训练损失lr该 epoch 的学习率dice coefficient验证集上的 Dice 系数验证集混淆矩阵逐类别交并比/像素准确率等统计。结合--save-best机制训练过程可以在「验证 Dice 提升即覆盖 best_model、否则跳过保存」的策略下自动保留最优权重便于后续预测或继续调参。九、注意事项与常见问题汇总 README 及源码中的关键注意点--data-path必须指向 DRIVE 根目录脚本内部会拼接DRIVE/training、DRIVE/test若路径下找不到DRIVEtrain_multi_GPU.py 会直接抛出异常预测时务必修改weights_path为实际生成的权重路径否则脚本会因找不到权重断言失败predict.py使用 validation 相关代码时确保验证集/测试集包含每个类别的目标只需修改--num-classes、--data-path和--weights其他代码尽量不要改动ROI 掩码语义标签中 255 表示 ROI 外区域在损失与评估中通过ignore_index255忽略训练时collate_fn也用 255 填充 batch三处语义保持一致上采样默认使用双线性插值若追求更丰富的可学习上采样可将bilinearFalse切换为转置卷积对应 src/unet.py 中的分支更换数据集时请用 compute_mean_std.py 重新统计归一化参数并按需调整--num-classes、损失权重与类别映射。十、扩展与进阶方向更换编码器主干仓库提供了 src/vgg_unet.py 与 src/mobilenet_unet.py 两个变体可作为在轻量化或更强特征提取能力之间权衡的起点损失函数调优在 train_and_eval.py 中二分类场景已默认给前景 2 倍损失权重对于更严重的类别不均衡可进一步调整loss_weight或替换 Dice 系数中的 epsilon 平滑策略与其他分割模型对比本仓库pytorch_segmentation目录下还包含 deeplab_v3、fcn、lraspp、u2net 等分割实现可基于同一数据集横向对比不同架构的分割精度与推理速度部署与转换如需将训练好的 U-Net 部署到服务端可参考仓库deploying_service目录下关于 ONNX/OpenVINO/TensorRT 转换的通用流程将best_model.pth导出为推理引擎支持的格式。通过本文的完整流程你可以基于本仓库从零训练一个用于视网膜血管分割的 U-Net 模型并掌握数据统计、增强对齐、组合损失、分布式训练与推理后处理等贯穿图像分割任务全生命周期的方法这些技能可直接迁移到其他二分类/多分类分割场景。赞分享示例工程【免费下载链接】deep-learning-for-image-processingdeep learning for image processing including classification and object-detection etc.项目地址https://gitcode.com/gh_mirrors/de/deep-learning-for-image-processing点击查看免费下载相关推荐ORVS 眼底视网膜血管分割实战基于 MMSegmentation 的 UNet 训练与评估指南ORVS 眼底视网膜血管分割实战基于 MMSegmentation 的 UNet 训练与评估指南 ORVSOnline Retinal image for人工智能深度学习计算机视觉终极指南基于预训练ResNet-50的U-Net图像分割实战教程终极指南基于预训练ResNet 50的U Net图像分割实战教程 在计算机视觉领域图像分割技术正经历着革命性的变革。面对日益增长的应用需求开发者们迫切需要深度学习计算机视觉5分钟快速搭建基于预训练ResNet-50的U-Net图像分割实战指南5分钟快速搭建基于预训练ResNet 50的U Net图像分割实战指南 在计算机视觉领域图像分割任务面临着训练时间长、数据需求大、模型泛化能力不足等核心挑战深度学习计算机视觉上一篇Hermes Agent 启动报错怎么办看懂四拍启动链4 步排障 10 分钟定位下一篇掌握ThinkPad散热控制TPFanControl2完全指南与静音优化方案创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考