Pytorch Unet医学图像分割实战:一键训练脚本与预测全流程解析
发布时间:2026/10/2 14:57:53来源:尧图网络
简介一个基于Pytorch与Unet的医学图像分割实战项目面向有一定深度学习基础的开发者、医学影像研究者以及需要快速落地分割任务的技术人员适用于病灶区域提取、器官结构分割等实际场景。项目完整覆盖数据加载、模型搭建、模型训练、验证评估与推理预测全流程采用UNet的对称收缩与扩展路径设计借助跳跃连接实现多尺度特征融合同时附带一键执行训练脚本配置好依赖后即可启动训练提前保存的模型权重文件使无GPU环境下也能直接运行预测。压缩包共99个文件整体大小约122MB文件类型包括90张图片样本、3个Python源码、1个Shell一键训练脚本、1个PyTorch权重文件另有说明文档与依赖清单目录划分清楚便于定位代码和结果。当前已有587人学习下载适合作为医学图像分割入门练习、毕业设计复现或算法二次开发的参考基础。1. 医学图像分割遇上一键训练脚本这套Unet实战项目到底值不值得跑医学图像分割这个方向落到实际操作上就是一件事让网络把CT、MRI或病理切片里的器官、病灶逐像素圈出来。Pytorch加Unet的组合之所以满大街都是不是因为大家偷懒而是对于几万张或几百张的医学数据Unet的参数量和收敛特性刚好踩在那个能跑、能改、能出结果的甜区。标题里说支持训练预测一键执行训练脚本这意味着你拿到的不是一堆散装代码而是能直接换成自己数据跑起来的项目骨架对刚入坑分割任务的人尤其友好。这套东西适合谁呢手里有标注好的影像数据但没跑过分割模型的科研人员、刚看完Pytorch基础想找个完整项目的学生、以及要快速验证某个器官分割可行性的算法工程师。下面按我自己落地这类项目的流程来拆为什么选Unet、训练脚本怎么组织、预测流程怎么对齐、哪些坑你大概率会踩。2. 用Pytorch重写Unet的落地选型为什么医学分割默认拿它当基线2.1 Unet的编码器-解码器结构在医学影像上的三个先天优势Unet名字来自它的U形结构左边一串卷积和下采样不断压缩特征图右边一串上采样把特征图恢复回原始分辨率中间用skip connection把编码器每一层的特征拼到解码器对应层。这种设计对医学影像几乎是量身定做的。第一个优势是浅层信息和深层信息都保留。医学图像里器官边界往往是模糊的对比度也低如果只用深层语义特征分割出来的边缘会像被橡皮擦蹭过。skip connection把浅层的高分辨率边缘特征直接传到解码器让网络在做像素分类的同时还能看见原始边界。第二个优势是小样本也能训练。医学分割数据集经常只有几十到几百张像ResNet这种几十层的分类骨架预训练权重又不好找。Unet的基础版本只有大约3100万参数显存占用也温和两三张标注图也能把它推向一个可用的局部最优。第三个优势是输入输出天然同尺寸。医学分割要求输出和输入一样大的maskUnet没有全连接层任何尺寸的输入都能得到对应尺寸的输出。训练和预测时只要把图片resize到网络输入尺寸不需要像分类网络那样对FeatureMap做全局池化。选择Pytorch的理由也很直接。医学分割需要频繁改网络结构和调试数据预处理Pytorch的动态图机制能让你在模型里直接打印中间层形状出错了在抛异常之前就能看出来。torchvision和MONAI这些生态也都在Pytorch这边后面想换成AttentionUnet或者加预训练encoder社区里现成的实现基本都能直接用。2.2 Pytorch做医学分割的理由与最小环境配置配置环境的流程里最常见也最省心的步骤是用Anaconda建独立虚拟环境。千万别图省事把torch装在base环境里项目依赖一旦升级其他项目的torch版本会被一起动掉。conda create -n medseg python3.9 conda activate medseg pip install torch2.0.1 torchvision0.15.1 --index-url https://download.pytorch.org/whl/cu118 pip install numpy opencv-python tqdm tensorboard这里需要说清楚两个细节。第二行的--index-url对应CUDA 11.8版本如果显卡驱动只支持CUDA 10.2就得换对应后缀。先确认NVIDIA驱动版本在命令行执行nvidia-smi看右上角CUDA Version一般不高于这个版本号就能用。第三行装的是医学分割最少需要的库numpy做数组运算、opencv读写图片、tqdm显示训练进度、tensorboard看loss曲线。环境装完别急着跑大网络先跑一个验证句子。import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0))输出True说明Pytorch正确调用了显卡。如果输出False先别怀疑代码回去查NVIDIA驱动和Pytorch的CUDA版本是否匹配。这一步解决了后面百分之九十的环境问题都不会再找你。2.3 医学分割数据集长什么样单通道灰度图同尺寸mask医学分割数据集和自然图像分割最大的差别是图像几乎全是单通道灰度图mask是像素级标注的索引图或二值图。以最常见的肝脏CT分割为例images文件夹里是CT的PNG灰度图masks文件夹里是对应名称的掩膜前景像素为255背景为0。目录结构一般是dataset/ images/ patient001.png patient002.png masks/ patient001.png patient002.png这里有个新手极易忽略的点mask和images的文件名必须一一对应而且图片尺寸一致。很多数据集是用ITK-SNAP从DICOM上手工标注导出的导出时可能只保存了病灶区域导致图像是512x512mask是128x128。训练脚本加载时对不上要么报错要么错位。预处理时我一般会写一个简单的数据确认函数先打印训练集和mask的形状、唯一值而不是等到训练开始后才被batch大小不一致的报错打断。import cv2 import numpy as np img cv2.imread(dataset/images/patient001.png, cv2.IMREAD_GRAYSCALE) mask cv2.imread(dataset/masks/patient001.png, cv2.IMREAD_GRAYSCALE) print(image:, img.shape, img.dtype, img.min(), img.max()) print(mask:, mask.shape, mask.dtype, np.unique(mask))这段代码输出里mask唯一的合理值是[0, 255]如果出现[0, 1, 255]说明标注把多个器官标成了不同值训练时记得把非零值统一改成1。这就是后面预测全黑白问题的根源提前确认能省很多事。3. 训练脚本怎么落地从项目文件到一键跑通的完整流程3.1 项目文件结构与训练入口设计一个能一键执行的Pytorch Unet项目文件结构应该清楚到不需要看README就能知道每个文件干什么。我常用的最小骨架是三文件一目录unet_seg/ model.py # Unet网络定义 dataset.py # 数据集加载与预处理 train.py # 训练主逻辑 predict.py # 预测主逻辑 config.py # 全局参数配置 checkpoints/ # 模型权重保存目录model.py只做一件事定义Unet类。dataset.py负责从文件夹读图片和mask并提供__getitem__方法。train.py是训练入口predict.py是预测入口。把参数集中放在config.py里比在train.py顶部写一堆常量更好维护因为训练和预测都要用到统计数据比如归一化的均值方差。train.py最外层的流程是读配置、初始化模型、初始化Dataset和DataLoader、定义损失函数和优化器、开始epoch循环。主循环内部每训练一个epoch就验证一次根据验证集Dice决定是否保存checkpoint。防止训练中途显存爆掉可以用torch.cuda.empty_cache()做兜底但别指望它解决根本问题。3.2 一键执行训练脚本的写法与启动顺序这个项目最值钱的部分是一键执行训练脚本。很多从GitHub下载的项目训练代码写得很完整但需要你手动激活环境、手动跑好几条命令才能启动。一键脚本的价值是把这些步骤固化下来跑错一步就停下来跑通一次以后永远一样。#!/bin/bash set -e source activate medseg cd $(dirname $0) if [ ! -d dataset/images ]; then echo error: dataset/images not found exit 1 fi python train.py --epochs 100 --batch_size 8 --lr 1e-4两行关键内容说明一下。set -e的意思是脚本中任一命令返回非零状态码就立即终止。没有它环境激活失败时脚本会继续跑trian.py最终抛出一堆torch报错而真正的问题——conda环境没激活——早被淹没了。cd $(dirname $0)是让脚本无论从哪里被调用都先切换到脚本所在目录。你从项目根目录运行bash train.sh和从其他路径加绝对路径运行行为完全一致。Windows下对应的train.bat长这样echo off call conda activate medseg cd /d %~dp0 python train.py --epochs 100 --batch_size 8 --lr 1e-4 pause%~dp0的功能和上面shell的dirname一致取当前脚本所在目录。注意batch里不加set -e但python train.py执行失败后pause能保证窗口不闪退你还能看到错误信息截图回去排查。3.3 训练参数怎么定batch size、学习率、早停与checkpoint训练参数按数据规模和显存两个维度来定。医学分割图普遍偏大常见的原始尺寸是512x512甚至1024x1024网络下采样四层再上采样回来后中间特征维度是原始图片的1/16。一个512x512的输入在16倍下采样层有512x32x32的通道数光这一层的feature map就大约占用512MB显存所以batch size不要拍脑袋填。我的一般规律是给数据集跑一次profiler先把batch_size设成2训练一个step看显存峰值如果显存占用低于总显存的60%再翻倍。学习率方面Unet这种每层卷积后面都跟着BatchNorm的网络初始学习率设置成3e-4通常起步很稳。如果发现训练到第10轮loss还在原地把学习率降到1e-4再试。早停和checkpoint是训练可靠性的关键。别用固定epoch数硬跑医学分割数据的验证集Dice往往在第30轮到第60轮之间出现抖动继续训练可能过拟合。用验证集Dice作为早停依据best_dice 0.0 patience 20 bad_epochs 0 for epoch in range(epochs): train_one_epoch(model, train_loader, optimizer) val_dice validate(model, val_loader) if val_dice best_dice: best_dice val_dice torch.save(model.state_dict(), checkpoints/best_model.pth) bad_epochs 0 else: bad_epochs 1 if bad_epochs patience: print(fearly stop at epoch {epoch}, best dice {best_dice:.4f}) break这段代码里最重要的不是保存逻辑而是bad_epochs只有验证集Dice没破纪录时才累加。医学分割验证集Dice经常出现连续30个epoch不涨但第31个epoch突然跳升的情况patience设20到30是合理的。保存模型尽量存state_dict而不是整个模型对象前者只包含权重换个文件结构也能加载后者把类定义也序列化进去改了模型结构就废了。4. 训练好之后怎么做预测单张图推理的完整流水线4.1 加载模型权重与预处理对齐预测代码和训练代码最大的不同是模型必须切换到eval模式并且关闭梯度计算。分三步走。第一步用和训练完全一样的模型定义去实例化网络结构。别在predict.py里重新写一个结构不一样的Unet哪怕是少个卷积层加载权重就会提示shape不匹配报错信息是size mismatch。第二步加载checkpoint。from model import Unet model Unet(n_channels1, n_classes1) state_dict torch.load(checkpoints/best_model.pth, map_locationcpu) model.load_state_dict(state_dict) model.to(device) model.eval()map_locationcpu这个参数是为了没有GPU的机器上也能加载如果你的预测机器和训练机器用的同一张显卡类型可以直接写map_locationcuda:0省去后续CPU到GPU的拷贝。第三步是对齐预处理。很多预测脚本翻车不是模型本身而是训练时做了(x - 0.5) / 0.5归一化预测时直接读原图喂进去。最常见的现象是预测mask一片白或者一片黑。我见过一个项目训练时把所有图片resize到256x256预测时忘了resize直接拿512x512图进去跑输出mask也是512x512但是边缘和器官完全错位。所以预处理必须抽成函数训练和预测都用同一个。4.2 推理、阈值分割与mask保存单张图的推理流程如下import cv2 import numpy as np import torch image cv2.imread(dataset/images/patient001.png, cv2.IMREAD_GRAYSCALE) image cv2.resize(image, (256, 256), interpolationcv2.INTER_LINEAR) image image.astype(np.float32) / 255.0 image (image - 0.5) / 0.5 input_tensor torch.from_numpy(image).unsqueeze(0).unsqueeze(0) input_tensor input_tensor.to(device) with torch.no_grad(): output model(input_tensor) output torch.sigmoid(output).squeeze().cpu().numpy() seg_mask (output 0.5).astype(np.uint8) * 255这一段有三个值得注意的地方。第一unsqueeze(0).unsqueeze(0)分别加了batch维度和channel维度因为模型接收的是NCHW的四维张量单张灰度图原本只有HW两维。第二with torch.no_grad()必须写否则推理时会为中间结果建立计算图显存占用直接翻两三倍。第三输出经过sigmoid变成0到1的概率图然后以0.5为阈值做二值化。阈值为什么是0.5不是0.7因为Unet的最后一层通常是一个卷积层输出logits经过sigmoid映射为概率。如果训练时用的损失是BCEWithLogitsLoss它在内部做了sigmoid预测时你也要手动做sigmoid。如果只取logit的符号结果等价于阈值0.5但概率图的语义不够直观不方便调节敏感度。保存mask和可视化结果cv2.imwrite(predictions/patient001_mask.png, seg_mask) color np.zeros((seg_mask.shape[0], seg_mask.shape[1], 3), dtypenp.uint8) color[:, :, 2] seg_mask overlay cv2.addWeighted(cv2.cvtColor(original_resized, cv2.COLOR_GRAY2BGR), 0.6, color, 0.4, 0) cv2.imwrite(predictions/patient001_overlay.png, overlay)把mask贴在原图上红色半透明叠加在器官区域这一步不是形式主义。医生和标注人员看叠加图能快速判断分割边界是否贴合解剖结构给自己看也能一眼发现mask是不是整块偏移了。4.3 批量预测与结果可视化单张图能跑通后批量预测只是加一层文件夹遍历。但要小心输出mask的大小应该和输入图片原始分辨一致。训练时的resize只是为了进网络预测完要把输出mask重新resize回原始尺寸否则保存的mask比原图小一圈后续做体积计算时数值全错。import os from tqdm import tqdm os.makedirs(predictions, exist_okTrue) for name in tqdm(os.listdir(dataset/images)): img_path os.path.join(dataset/images, name) mask_path os.path.join(predictions, name.replace(.png, _mask.png)) image cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) h, w image.shape[:2] resized cv2.resize(image, (256, 256)) input_tensor torch.from_numpy(resized.astype(np.float32) / 255.0) input_tensor input_tensor.unsqueeze(0).unsqueeze(0).to(device) with torch.no_grad(): output model(input_tensor) seg (torch.sigmoid(output).squeeze().cpu().numpy() 0.5).astype(np.uint8) mask_full cv2.resize(seg, (w, h), interpolationcv2.INTER_NEAREST) cv2.imwrite(mask_path, mask_full * 255)注意倒数第三行mask回缩到原始尺寸时必须用INTER_NEAREST最近邻插值。如果用线性插值0和1之间会出现0.4、0.6这种中间值导致保存的PNG出现锯齿或灰色边缘。这是预测一个极其隐蔽的坑值得单独记录。5. 医学图像分割常见问题排查三个必踩的坑与我的处理方式5.1 训练loss不降标签与类别权重先查这两处现象训练跑了十几个epochloss一直徘徊在0.7附近不下去验证集Dice几乎为零。原因最常见的是标签类别严重不平衡。医学分割里目标区域往往只占整张图的5%以下比如肺结节分割背景像素占绝大比例。BCE损失对每个像素是完全平等的网络发现把所有像素预测成背景就能拿到一个低得离谱的loss于是模型很快就塌缩到全背景。解决先打印一个batch里mask的前景像素占比。如果低于10%把损失函数从BCEWithLogitsLoss换成DiceLoss或BCE Dice的组合。Dice损失天然按类别交集占比计算不在乎前景区域小。另一个办法是给BCE加权pos_weight torch.tensor([background_pixels / foreground_pixels]) criterion torch.nn.BCEWithLogitsLoss(pos_weightpos_weight)pos_weight设为背景像素数除以前景像素数相当于对前景像素的错误给更大惩罚。这个方法只在两类分割时好使多类分割还得换Focal Loss。5.2 预测mask全黑或全白归一化与resize的玄学现象训练时Dice有0.8但预测单张图存出来的mask是全黑的。原因大概率是预测代码里的预处理和训练不一致。遇到过两种情况。第一种是训练用了(x - mean) / std标准差归一化mean和std是用整个训练集统计出来的预测时图省事只除以255。第二种是mask在resize时用了线性插值预测输出概率图直接resize中间值的0.4、0.2被阈值0.5一卡全部归零结果看起来就全黑了。解决把预处理和resize封装到同一个函数训练循环和预测脚本都调用它。mask的回缩必须用最近邻插值这条规律在上面批量预测的代码里已经写过再踩一次不值得。5.3 GPU显存溢出Unet虽然轻batch再小也会翻车现象用512x512输入训练batch_size设成8刚跑完第一个batch就报RuntimeError: CUDA out of memory。原因很多人以为Unet参数量才3100万显存很宽裕忽略了中间层的feature map。实际上Unet前半部分每层卷积输出的通道数分别是64、128、256、512对应分辨率减半。一个512x512的输入第一层feature map就是64x512x512单个batch占16MB8个batch约128MB再加上反向传播保存的中间梯度实际显存占用大约是显式计算的2到3倍。解决第一优先把batch_size降到4或2第二优先用AMP混合精度训练。Pytorch的自动混合精度能让显存占用降低40%甚至更多。scaler torch.cuda.amp.GradScaler() for batch in train_loader: with torch.cuda.amp.autocast(): outputs model(images) loss criterion(outputs, masks) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()用AMP后特征图以半精度存储显存减半batch从4提到8不是梦。但它也有副作用BatchNorm在混合精度下统计量可能不稳定初期训练建议前10个epoch用全精度后面再开AMP。5.4 数据加载慢到怀疑人生IO才是分割项目的隐藏瓶颈现象GPU利用率只有20%到30%CPU风扇狂转训练一个epoch要十分钟以上。原因训练循环本身没瓶颈瓶颈在每次迭代时从硬盘读取512x512的PNG并解码。医学分割数据量大、图像大单读盘就占满一个核显存里GPU早就算完了等CPU喂下一批。解决把DataLoader的num_workers从默认0改为CPU核数的一半并开启pin_memoryTrue。另一个更有效的方法是训练前把所有图片预处理成内存npy文件一次性读取。from torch.utils.data import Dataset, DataLoader class MedSegDataset(Dataset): def __init__(self, preprocessed_dir): self.images np.load(os.path.join(preprocessed_dir, images.npy)) self.masks np.load(os.path.join(preprocessed_dir, masks.npy)) def __getitem__(self, idx): img torch.from_numpy(self.images[idx]).float() mask torch.from_numpy(self.masks[idx]).float() return img.unsqueeze(0), mask.unsqueeze(0) train_loader DataLoader(dataset, batch_size8, num_workers4, pin_memoryTrue)如果数据集太大一次性读入内存超过32GB就退一步把图片在首次epoch时缓存为PNG的等尺寸压缩格式。从内存读npy比从硬盘读PNG快5到8倍GPU利用率从20%拉到70%是常见结果。5.5 训练与预测的图片尺寸不一致错误跑到最后一步才发现现象训练时输入固定resize到256x256预测时没resize输出mask只有原始尺寸的1/4但程序不报错直到把mask叠加到原图上才发现对不齐。原因Unet只有卷积层理论上能处理任意尺寸输入所以训练时resize的固定尺寸被预测代码忽略了。没有全连接层的网络就是这么宽容错了也不提醒。解决在预测脚本里把固定尺寸和插值方式写进一个常量用断言保护assert input_tensor.shape[-2:] (256, 256), input size must match training size预测代码一旦跑出shape不一致的mask直接中断不要带病输出。懒得多写一行的后果是后面所有统计分析都得重跑。6. 让分割结果更可信Dice评估、数据增强与模型改进的三个方向训练脚本能跑通只是起点医学分割的价值在评估指标和应用侧。我验证模型时最先看的不是像素准确率而是Dice系数和边界距离。Dice的计算方式在二值分割里是两个mask交叠面积的2倍除以两者面积之和。对于小器官靠像素准确性很容易产生背景预测正确率高但是病灶一个没圈出来的假象Dice对漏检和误检同样敏感。在验证集上我还会算一下预测mask的连通域数量如果出现碎片化的小岛说明模型对纹理敏感但对边界保守通常需要给损失函数加一项边界惩罚。数据增强的方向也和自然图像不同。随机翻转、旋转90度这类几何增强能提升模型对体位变化的鲁棒性我在训练时还会加入弹性形变模拟器官在呼吸运动下的形变。强度上别加高斯噪声和亮度抖动医学图像灰度级本来就受设备和协议影响增强做得太狠模型会学坏。模型本身的改进我推荐三个方向按难度排序第一是在解码器每一层上采样前拼上对应的编码器特征图并做一个3x3卷积再拼接这个操作对边界精度的提升立竿见影。第二是引入注意力门控让网络自发忽略背景区域的响应。第三是如果数据量在万张以上把编码器换成ResNet34预训练权重再用torchvision提供的模型微调。最后提醒一句我每次训练前先订好固定随机种子记录所有超参数这样翻车了还能改回去。希望帮到你。本文还有配套的精品资源点击获取
网站建设高端定制企业官网