ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

UNet与UNet++细胞分割源码实战:从训练到SAHI切片预测

UNet与UNet++细胞分割源码实战:从训练到SAHI切片预测 简介这份源码面向计算机相关专业的毕业设计、课程设计与期末综合作业场景提供基于UNet与UNet两种编码器-解码器架构的细胞医学图像分割Python实现适合需要项目实践训练或算法复现的学习者。压缩包共58个文件约107KB以44个py脚本为核心覆盖数据加载、数据增强、模型构建、训练与预测全流程另含zbak备份、Dockerfile、requirements.txt与readme.md等环境配置与说明文件便于快速搭建可复现实验环境。项目原为本科课程设计在导师指导下完成并获99分评价代码结构完整、注释详细包含损失函数配置、dice_score等评估指标以及UNet与UNet的对比实验能帮助读者理解两种网络在医学图像处理中的特性差异。目前已有62人学习适合作为分割任务入门与进阶的参考实现。1. 从一份 99 分的课设说起UNet 与 UNet 细胞分割源码能跑出什么如果你正在做医学图像分割相关的毕业设计或期末大作业大概率绕不开 UNet 这个经典结构。但网上能找到的代码要么只有模型定义没有训练流程要么跑起来就报维度不匹配要么评估指标写得含糊其辞。这份基于 UNet 与 UNet 的细胞医学图像分割 Python 源码是我近期拆过结构比较完整的一份——它把数据加载、模型构建、训练循环、Dice 评估、切片预测全串起来了还额外带了一套 SAHI 切片推理工具链。它解决的核心问题是让你不用从零搭训练框架直接在一份能跑通的代码上理解编码器-解码器怎么落地到细胞分割任务。适合两类人一是计算机相关专业需要交课程设计或毕设的学生二是想快速验证 UNet 系列在自己数据上表现的开发者。代码里 UNet 和 UNet 两条路线都有方便你做对比实验。下面我按实际拆包顺序把关键模块、参数配置和踩过的坑逐一讲清。2. 环境搭建与数据准备requirements 里没写全的依赖怎么补2.1 先看清目录结构再动手装环境拿到压缩包解压后根目录下有这么几个关键文件train.py是训练入口predict.py和slicePredict.py分别对应整图预测和切片预测evaluate.py负责指标计算requirements.txt列了基础依赖Dockerfile给了容器化方案。模型定义在unet/目录下unet_model.py是主结构unet_parts.py封装了卷积块和下采样、上采样组件。utils/里放的是数据加载和 Dice 计算sahi/是一套切片推理框架scripts/下还有 COCO 格式转换和 FiftyOne 可视化脚本。这个结构比很多课设代码规范但requirements.txt我打开看了一眼只写了 torch、numpy、opencv-python 这几个大件实际跑起来还缺一些。常见做法是先建虚拟环境再逐个补。# 创建虚拟环境Python 3.8 兼容性最稳 python -m venv venv_unet source venv_unet/bin/activate # Windows 用 venv_unet\Scripts\activate # 先装 requirements 里的基础包 pip install -r requirements.txt # 补装实际运行需要的包 pip install scikit-image pillow tqdm tensorboard pip install sahi fiftyone # 如果要跑切片预测和可视化脚本这里有个参数要注意torch版本建议选 1.10 到 1.13 之间太新的 2.x 版本在unet_parts.py里某些nn.Module的初始化写法上可能报 warning。如果你用 GPU去 PyTorch 官网查对应 CUDA 版本的安装命令别直接pip install torch装成 CPU 版。2.2 数据目录怎么摆、增强参数在哪调代码默认的数据加载逻辑在utils/data_loading.py里它期望的数据结构是训练集和验证集分开每张图像对应一个 mask。我一般会整理成这样的目录data/ train/ images/ *.png masks/ *.png val/ images/ *.png masks/ *.pngmask 必须是单通道灰度图细胞区域为白色像素值 255背景为黑色。如果你拿到的原始数据是彩色标注图需要先转成二值 mask否则dice_score.py算出来的值会偏低。数据增强策略在data_loading.py里通过albumentations库实现常见的有随机旋转、水平垂直翻转、弹性变形。弹性变形对细胞图像特别有用因为细胞形态本身不规则。我一般会把弹性变形的alpha参数设在 120 左右sigma设在 12 左右太大会把细胞拉变形到失去语义。# data_loading.py 里增强部分的典型配置 import albumentations as A train_transform A.Compose([ A.RandomRotate90(p0.5), A.HorizontalFlip(p0.5), A.VerticalFlip(p0.5), A.ElasticTransform(alpha120, sigma12, p0.3), A.Resize(height256, width256), # 统一尺寸UNet 要求输入固定 ])Resize的尺寸要和train.py里的--img_size参数一致默认是 256。如果你显存够可以调到 512但 batch size 要相应降到 4 或 2。细胞分割任务里输入尺寸太小会丢失小细胞太大又吃显存256 是个折中起点。3. UNet 与 UNet 模型实现编码器-解码器到底怎么搭3.1 UNet 的下采样与跳跃连接在代码里长什么样unet_parts.py里定义了两个核心类DoubleConv和Down、Up。DoubleConv就是两次Conv2d BatchNorm ReLU的堆叠这是 UNet 的基本计算单元。Down负责下采样先 maxpool 再接DoubleConvUp负责上采样用ConvTranspose2d把特征图放大后和编码器对应层的特征做通道维度拼接。# unet_parts.py 里 Up 模块的关键逻辑 class Up(nn.Module): def __init__(self, in_channels, out_channels, bilinearTrue): super().__init__() if bilinear: self.up nn.Upsample(scale_factor2, modebilinear, align_cornersTrue) else: self.up nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size2, stride2) self.conv DoubleConv(in_channels, out_channels) def forward(self, x1, x2): x1 self.up(x1) # 处理尺寸不整除时的 padding diffY x2.size()[2] - x1.size()[2] diffX x2.size()[3] - x1.size()[3] x1 F.pad(x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2]) x torch.cat([x2, x1], dim1) # 跳跃连接编码器特征 解码器特征 return self.conv(x)这里bilinearTrue时用双线性插值上采样参数量少False时用转置卷积效果可能略好但容易产生棋盘伪影。细胞分割任务里我一般先用双线性插值跑 baseline确认流程通了再换转置卷积对比。跳跃连接是 UNet 的灵魂它把编码器的高分辨率特征直接送到解码器弥补下采样过程中的空间信息丢失。但要注意拼接时通道数会翻倍所以DoubleConv的输入通道要设对否则会报size mismatch。3.2 UNet 的嵌套密集连接改了什么UNet 的核心改动是在编码器和解码器之间加了一系列嵌套的卷积层形成密集跳跃连接。unet_model.py里如果实现了 UNet你会看到类似NestedUNet的类它的forward里维护了一个多层特征列表每一层解码器都会接收前面所有同分辨率层和上一层上采样的输出。# UNet 嵌套连接的简化示意 # 假设编码器输出为 x0_0, x1_0, x2_0, x3_0, x4_0 # 解码器节点 x0_1 接收 [x0_0, up(x1_0)] # 解码器节点 x0_2 接收 [x0_0, x0_1, up(x1_1)] # 以此类推每个节点都融合了前面所有同尺度信息这种设计的好处是梯度流动更充分对小目标分割更友好。代价是参数量和显存占用比 UNet 高不少。我在 256x256 输入、batch size 8 的设置下UNet 大概占 3GB 显存UNet 要到 5GB 左右。如果你的显卡只有 6GB跑 UNet 时把 batch size 降到 4 或者用梯度累积。选型建议如果细胞边界清晰、大小均匀UNet 够用如果细胞大小差异大、有粘连UNet 的密集连接通常能带来 2 到 5 个百分点的 Dice 提升。但别盲目上 UNet先跑通 UNet 确认数据和流程没问题再换模型对比。4. 训练、评估与切片预测从 train.py 到 slicePredict.py 的完整链路4.1 损失函数与优化器参数怎么设train.py里默认用的是BCEWithLogitsLoss加DiceLoss的组合这是医学分割里很常见的搭配。BCE 负责像素级分类Dice 负责优化区域重叠度。代码里一般会写成criterion lambda pred, target: bce(pred, target) dice(pred, target)的形式。# 启动训练的命令示例 python train.py \ --data_dir ./data \ --model unet \ --epochs 100 \ --batch_size 8 \ --lr 1e-3 \ --img_size 256 \ --val_percent 0.1 \ --save_checkpoint参数说明--model可以选unet或unet--lr初始学习率设 1e-3如果 loss 震荡明显降到 1e-4--val_percent是从训练集里划多少做验证0.1 表示 10%。--save_checkpoint会保存每个 epoch 的最优模型存在checkpoints/目录下。优化器默认是 Adam动量参数用默认的 0.9 和 0.999 就行。学习率调度我一般加一个ReduceLROnPlateau当验证 Dice 连续 5 个 epoch 不升就乘 0.5。这个在train.py里可能没写需要自己补几行。4.2 Dice 评估与切片预测的实操细节evaluate.py算的是 Dice 系数和 IoU输出格式一般是每个验证样本的分数加平均值。Dice 的计算逻辑在utils/dice_score.py里核心就是2 * 交集 / (预测和 真实和)。注意如果预测和真实都为空Dice 定义为 1代码里要处理这个边界否则会除零。# dice_score.py 的核心计算 def dice_coeff(pred, target, smooth1e-6): pred pred.contiguous().view(-1) target target.contiguous().view(-1) intersection (pred * target).sum() return (2. * intersection smooth) / (pred.sum() target.sum() smooth)smooth参数是防止除零的别省。评估时要把模型设成eval()模式并关掉梯度否则 BatchNorm 的统计量会变Dice 会偏低。slicePredict.py是这套代码里比较有特色的部分它用 SAHI 框架把大图切成小块分别预测再拼回去。这对病理切片这种超大分辨率图像很有用因为整图直接 resize 到 256 会丢失大量细节。切片预测的关键参数是slice_height和slice_width一般设成和训练输入一致overlap_ratio设 0.2 到 0.3避免拼接处出现断裂。# 切片预测命令示例 python slicePredict.py \ --model_path checkpoints/best_model.pth \ --source ./test_images \ --slice_size 256 \ --overlap 0.25 \ --output ./predictionsoverlap太大会增加推理时间太小会在切片边界漏掉细胞。0.25 是我试过比较平衡的值。预测结果会保存成 mask 图可以用scripts/下的可视化脚本叠加到原图上检查。5. 避坑与排查这份代码跑起来最容易翻车的五个地方5.1 报错 “Expected 4D input” 或维度不匹配现象训练刚启动就报RuntimeError: Expected 4D (got 3D) input to Conv2d。原因通常是数据加载时没有加 batch 维度或者 mask 的通道数和预测输出对不上。解决检查DataLoader的batch_size是否大于 1检查 mask 是否被读成了三通道。在data_loading.py里加一句mask mask.convert(L)强制转灰度。5.2 Dice 一直卡在 0.3 左右不上升现象训练 loss 在降但验证 Dice 不动。原因可能是学习率太大导致模型在局部震荡或者数据增强太激进把细胞形态破坏了。解决先把学习率降到 1e-4把弹性变形的概率从 0.3 降到 0.1观察 10 个 epoch。如果还不行检查 mask 的像素值是不是 0 和 1 而不是 0 和 255代码里如果没做归一化255 会导致 loss 爆炸。5.3 显存溢出 “CUDA out of memory”现象跑 UNet 时 batch size 设 8 直接 OOM。原因UNet 的嵌套连接保留了更多中间特征图显存占用比 UNet 高 60% 以上。解决把 batch size 降到 4或者用torch.cuda.amp做混合精度训练。在train.py里加scaler torch.cuda.amp.GradScaler()和with torch.cuda.amp.autocast():能省将近一半显存。5.4 切片预测结果拼接处有断裂现象slicePredict.py输出的 mask 在切片边界处细胞被切断。原因overlap设得太小或者后处理时没有做融合。解决把overlap从 0.1 提到 0.25 以上检查sahi/postprocess/combine.py里的拼接逻辑是否用了加权平均而不是直接覆盖。如果代码里是直接覆盖手动改成按距离加权融合。5.5 评估指标比训练时低很多现象训练日志里 Dice 0.85跑evaluate.py只有 0.7。原因训练时用了数据增强评估时没关或者模型没设eval()模式BatchNorm 还在用 batch 统计量。解决在evaluate.py里确认model.eval()和torch.no_grad()都加了数据加载的 transform 只保留Resize和Normalize去掉所有随机增强。6. 进阶技巧用 FiftyOne 做分割结果可视化与错误分析跑通训练和预测之后真正花时间的是分析模型在哪里出错。scripts/目录下有一套 FiftyOne 集成脚本predict_fiftyone.py和coco2fiftyone.py可以把预测结果导入 FiftyOne 做交互式查看。FiftyOne 是个开源的可视化工具能让你按 Dice 分数排序、筛选低分样本、对比预测和真实 mask 的差异区域。我一般会先把预测结果转成 COCO 格式再用coco2fiftyone.py导入。转换脚本在scripts/coco2yolov5.py和scripts/coco_evaluation.py里也有参考。导入后可以按dice 0.5筛出失败案例看看是细胞粘连导致分割不全还是小细胞被漏掉。# 把预测结果导入 FiftyOne 做可视化 python scripts/predict_fiftyone.py \ --dataset_dir ./data/val \ --pred_dir ./predictions \ --model_path checkpoints/best_model.pth # 启动 FiftyOne 界面 fiftyone app launch在界面里可以叠加显示原图、真实 mask、预测 mask 三层用不同颜色区分真阳性、假阳性、假阴性。假阳性多说明模型把背景误判成细胞可能需要加负样本假阴性多说明小细胞被漏掉可以试试 UNet 或者把输入尺寸从 256 提到 384。还有一个实用技巧把evaluate.py输出的每个样本 Dice 存成 CSV用 pandas 按分数排序挑最低的 20 张单独看。这些样本往往暴露了数据标注问题——比如某些 mask 标注不一致或者图像本身模糊。模型表现上不去有时候不是网络结构的问题是标注质量的问题。从那以后我每次跑完训练都强制走一遍 FiftyOne 可视化再决定下一步调参方向光看平均 Dice 很容易被蒙蔽。希望这份拆解能帮你少走点弯路把这份源码真正跑成自己的东西。本文还有配套的精品资源点击获取
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进