拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

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

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

【免费下载链接】deep-learning-for-image-processing

deep 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,运行本项目建议满足以下环境:

  • Python:3.6 / 3.7 / 3.8;
  • PyTorch:1.10 及以上(训练脚本中大量使用torch.cuda.amp混合精度接口,请确保所用版本支持);
  • 操作系统:Ubuntu 或 CentOS(Windows 暂不支持多 GPU 训练,单卡/CPU 训练可用);
  • 硬件:最好使用 GPU 训练;多卡训练通过torchrun(PyTorch 1.10 起的官方分布式启动器)拉起多进程。

具体的依赖版本以 requirements.txt 为准,核心依赖如下:

numpy==1.22.0 torch==1.13.1 torchvision==0.11.1 Pillow

说明:requirements.txt中torch==1.13.1与torchvision==0.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四个基础模块拼装出完整网络:

  • DoubleConv(src/unet.py):连续两个Conv2d(3×3, padding=1, bias=False) + BatchNorm2d + ReLU(inplace=True)的序列,是编码器/解码器中的基本特征变换单元;
  • Down(src/unet.py):MaxPool2d(2, stride=2)后接DoubleConv,对应一次下采样;
  • Up(src/unet.py):当bilinear=True时使用nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),通道减半后经DoubleConv;当bilinear=False时改用nn.ConvTranspose2d(in_channels, in_channels//2, kernel_size=2, stride=2)转置卷积上采样。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], dim=1)

当输入宽高不是 2 的整数次幂时,编码器下采样会产生尺寸取整误差,该段代码通过中心式 padding 将上采样结果补齐到与跳跃连接特征一致,保证torch.cat拼接合法;

  • OutConv(src/unet.py):单个 1×1 卷积,将通道数映射为num_classes;
  • UNet主类(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_channels=3, num_classes=num_classes, base_c=32),即输入三通道 RGB、基础通道数为 32(显存占用更小,适合 DRIVE 这种小分辨率任务)。

值得一提的还有 src/vgg_unet.py 与 src/mobilenet_unet.py 两个变体文件,分别以 VGG、MobileNet 作为编码器主干替换默认的双卷积单元,为需要更换 backbone 的读者提供了扩展方向,本文默认介绍的是标准 U-Net 实现。

三、DRIVE 数据集准备与归一化统计

3.1 数据集下载与目录结构

DRIVE(Digital 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.gif)

3.2 数据集读取逻辑

my_dataset.py 中的DriveDataset实现了自定义Dataset:

  • 根据train=True/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_fn(my_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_size=int(0.5*base_size), max_size=int(1.2*base_size))、概率各为 0.5 的随机水平翻转与垂直翻转、RandomCrop(crop_size)、ToTensor、Normalize。默认base_size=565、crop_size=480(见 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 函数的执行逻辑可以概括为:

  1. 确定设备(cuda不可用则回退cpu);
  2. num_classes = args.num_classes + 1(背景 + 前景,DRIVE 场景下总类别为 2);
  3. 构建训练/验证DriveDataset与 DataLoader,其中num_workers取min(os.cpu_count(), batch_size, 8);
  4. 创建UNet(in_channels=3, num_classes=num_classes, base_c=32);
  5. 使用 SGD 优化器(momentum=0.9、weight_decay=1e-4);
  6. 创建混合精度GradScaler(仅当--amp开启时);
  7. 创建按 step 更新的学习率调度器(create_lr_scheduler,非按 epoch);
  8. 若指定--resume,加载 checkpoint 中的模型、优化器、调度器状态与start_epoch(AMP 时还会恢复 scaler);
  9. 逐 epoch 调用train_one_epoch与evaluate,将 loss、lr、Dice 系数写入resultsYYYYMMDD-HHMMSS.txt;
  10. 按--save-best保存save_weights/best_model.pth(Dice 最高)或save_weights/model_{epoch}.pth(每轮都存)。

5.3 损失函数与评估指标

训练与评估的核心实现位于 train_utils/train_and_eval.py:

  • criterion(train_and_eval.py):采用CrossEntropyLoss + DiceLoss的组合损失。交叉熵通过ignore_index=255忽略 ROI 外像素;当num_classes == 2时(背景/前景二分类),还会传入loss_weight=[1.0, 2.0](见 train_and_eval.py)加大前景(血管)在损失中的权重,缓解正负样本极度不均衡的问题;
  • DiceLoss(dice_coefficient_loss.py):对 softmax 后的概率图与 one-hot 标签计算 Dice 系数并取1 - dice作为损失,其中build_target将忽略像素在 one-hot 后仍标记为ignore_index,dice_coeff对忽略像素进行掩码剔除(dice_coefficient_loss.py);
  • evaluate(train_and_eval.py):以ignore_index=255构建混淆矩阵(ConfusionMatrix)与DiceCoefficient,在torch.no_grad()下逐 batch 更新指标,最后打印混淆矩阵与 Dice 系数;
  • create_lr_scheduler(train_and_eval.py):实现 warmup + poly 式学习率衰减。训练开始时倍率因子从warmup_factor=1e-3线性升至 1(1 个 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_node=8 train_multi_GPU.py

--nproc_per_node为使用的 GPU 数量。如需指定具体 GPU 设备,可在指令前添加CUDA_VISIBLE_DEVICES,例如只用物理设备中的第 1 块和第 4 块:

CUDA_VISIBLE_DEVICES=0,3 torchrun --nproc_per_node=2 train_multi_GPU.py

6.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 建议学习率乘以 n);
  • README 提醒Windows 暂不支持多 GPU 训练,分布式训练请使用 Linux。

七、模型预测与推理

训练完成后,可使用 predict.py 对测试图像进行推理。脚本核心流程如下:

  1. 配置输入:将weights_path设置为你训练生成的权重路径(默认./save_weights/best_model.pth),img_path指向测试图像(默认./DRIVE/test/images/01_test.tif),roi_mask_path指向对应的 ROI 掩码;
  2. 加载模型:UNet(in_channels=3, num_classes=classes+1, base_c=32),其中classes = 1(不含背景),并从 checkpoint 中取出'model'键加载权重(checkpoint 还包含优化器、调度器等训练状态);
  3. 预处理:使用与训练一致的mean/std做ToTensor + Normalize,并增加 batch 维度;
  4. 推理:model.eval()后先向网络送入一张与输入等尺寸的零张量(init_img)完成一次前向,再对真实输入计时推理;这一步常见于显存按需分配的 GPU 环境,可避免首次前向时的 CUDA 内核初始化时间计入推理耗时;
  5. 后处理:取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 及源码中的关键注意点:

  1. --data-path必须指向 DRIVE 根目录,脚本内部会拼接DRIVE/training、DRIVE/test;若路径下找不到DRIVE,train_multi_GPU.py 会直接抛出异常;
  2. 预测时务必修改weights_path为实际生成的权重路径,否则脚本会因找不到权重断言失败(predict.py);
  3. 使用 validation 相关代码时,确保验证集/测试集包含每个类别的目标;只需修改--num-classes、--data-path和--weights,其他代码尽量不要改动;
  4. ROI 掩码语义:标签中 255 表示 ROI 外区域,在损失与评估中通过ignore_index=255忽略,训练时collate_fn也用 255 填充 batch,三处语义保持一致;
  5. 上采样默认使用双线性插值,若追求更丰富的可学习上采样,可将bilinear=False切换为转置卷积(对应 src/unet.py 中的分支);
  6. 更换数据集时,请用 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-processing

deep learning for image processing including classification and object-detection etc.

项目地址:https://gitcode.com/gh_mirrors/de/deep-learning-for-image-processing
点击查看免费下载

相关推荐

上一篇:Hermes Agent 启动报错怎么办:看懂四拍启动链,4 步排障 10 分钟定位
下一篇:掌握ThinkPad散热控制:TPFanControl2完全指南与静音优化方案

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

返回列表