☰
基于UNet与UNet++的细胞图像分割实战:从源码解析到SAHI切片推理与调优
2026/10/5 4:08:21 网站建设 项目流程

简介:这份源码包面向计算机相关专业正在做毕业设计、课程设计或期末大作业的学生,以及需要医学图像分割实战练习的学习者,核心解决细胞图像分割任务的完整实现问题。项目基于UNet与UNet++两种经典网络结构,配套训练、评估与预测全流程代码,经导师指导并获评审99分,代码完整可直接运行,零基础也能上手。压缩包共48个文件,以44个Python源码为主,另含Dockerfile、requirements.txt、readme.md及.gitignore等配置说明文件,整体约95KB,体积轻量便于快速部署。目录中涵盖模型定义、数据加载、Dice评分、切片预测与后处理等模块,结构清晰,方便读者理解医学图像分割的工程组织方式。目前已有182人学习下载,适合作为毕设参考或课程实战模板,帮助读者掌握从数据预处理到模型推理的完整链路。

1. 从一份细胞分割源码包说起:UNet 与 UNet++ 到底能跑出什么

细胞图像分割这件事,真正上手做过的都知道,难点从来不是把网络搭出来,而是让模型在边界粘连、细胞重叠、染色不均的图像上还能把每个细胞分开。这份基于 UNet 和 UNet++ 的 Python 源码包,解决的就是这个场景:输入一张显微镜下的细胞图像,输出每个细胞的分割掩膜,用于后续计数、面积统计或形态分析。它适合正在做计算机相关专业毕业设计、课程设计、期末大作业的学生,也适合想拿一个完整医学图像分割项目练手的学习者。包里同时给了 UNet 和 UNet++ 两套模型,还带了 SAHI 切片推理、Dice 评分、数据加载和训练脚本,不是那种只丢一个模型文件让你自己猜怎么用的半成品。下面我按实际拆包和跑通的顺序,把这份资源讲清楚。

2. 拆开源码包:目录结构、模型选型与数据流

2.1 目录里到底有什么,哪些是核心

拿到压缩包解压后,根目录下能看到这些关键文件:train.py、evaluate.py、predict.py、slicePredict.py,以及unet/、sahi/、utils/、scripts/几个目录。requirements.txt和Dockerfile也在,说明作者考虑过环境复现。我一般先看unet/目录,因为模型定义全在这里。

unet/下有unet_model.py、unet_parts.py、__init__.py。unet_parts.py里通常是 DoubleConv、Down、Up、OutConv 这些基础模块,unet_model.py里则是 UNet 和 UNet++ 两个类的定义。UNet++ 的核心在于嵌套的密集跳跃连接和深监督,它比原版 UNet 多了一层解码器子网络,能缓解编码器与解码器之间语义鸿沟的问题。对于细胞图像这种目标小、边界模糊的场景,UNet++ 在边界召回上通常比 UNet 更稳,但参数量和显存占用也更高。

sahi/目录是这份包比较有诚意的地方。SAHI 是 Slicing Aided Hyper Inference 的缩写,专门做大幅图像的分块推理。细胞图像往往分辨率很高,直接整图送进网络会爆显存,或者下采样后小细胞直接消失。sahi/slicing.py负责切图,sahi/postprocess.py负责把各块的预测拼回原图,sahi/predict.py和sahi/prediction.py是推理入口。slicePredict.py就是调用这套流程的脚本。

utils/下有data_loading.py、dice_score.py、dataprocess.py、utils.py。data_loading.py管数据集读取和增强,dice_score.py实现 Dice 系数计算,dataprocess.py大概率是数据预处理或格式转换。scripts/里可能有辅助脚本,cli.py和annotation.py也在根目录附近,前者可能是命令行入口,后者可能涉及标注处理。

2.2 为什么选 UNet 和 UNet++,而不是 Transformer 分割

医学图像分割这几年确实有很多新架构,比如 Swin-UNet、TransUNet,但对细胞图像这个具体任务,UNet 系列仍然是很务实的选择。原因有三点:第一,细胞分割的训练数据通常不大,几百到几千张图,Transformer 类模型在小数据上容易过拟合,需要更强的预训练权重和更重的数据增强;第二,UNet 的跳跃连接对恢复细胞边界很有效,编码器下采样丢掉的细节,解码器能通过跳跃连接拿回来;第三,这份源码包已经把训练、评估、推理、切片推理都串好了,换成 Transformer 你得自己改数据流和损失函数,对做毕设或课程设计来说,时间成本不划算。

UNet++ 相比 UNet 的改进,主要是在跳跃连接上做了嵌套。原版 UNet 是编码器第 i 层直接连到解码器第 i 层,UNet++ 则在中间加了一系列密集卷积块,让不同尺度的特征先融合再往上传。这样做的好处是,解码器拿到的特征既有浅层的边缘信息,也有深层的语义信息,对细胞这种边界模糊的目标更友好。代价是显存占用增加,训练时间变长。如果你的显卡只有 6GB 或 8GB,建议先用 UNet 跑通,再换 UNet++ 对比。

2.3 数据从哪来,格式怎么组织

源码包本身不包含数据集,这是常见做法,因为医学数据涉及隐私和版权。你需要自己准备细胞图像数据集,常见的有 ISBI 2012 细胞分割数据集、DSB2018 数据科学碗数据集,或者自己标注的数据。数据组织方式一般是训练集和验证集分开,每张图像对应一个掩膜,掩膜是二值图,细胞区域为 1,背景为 0。

utils/data_loading.py里通常会定义一个 Dataset 类,读取图像和掩膜,做归一化、随机翻转、旋转等增强。我一般会先检查这个文件里的__getitem__返回的 tensor 形状和取值范围,确保图像是[C, H, W],掩膜是[1, H, W]或[H, W],值在 0 到 1 之间。如果这里不对,后面训练 loss 不降或者预测全黑,都是它引起的。

# 常见的数据加载类结构,具体以包内 data_loading.py 为准 class CellDataset(Dataset): def __init__(self, image_dir, mask_dir, transform=None): self.image_dir = image_dir self.mask_dir = mask_dir self.transform = transform self.images = os.listdir(image_dir) def __getitem__(self, index): img_path = os.path.join(self.image_dir, self.images[index]) mask_path = os.path.join(self.mask_dir, self.images[index].replace('.png', '_mask.png')) image = np.array(Image.open(img_path).convert('RGB')) mask = np.array(Image.open(mask_path).convert('L'), dtype=np.float32) mask[mask > 127] = 1.0 # 二值化,阈值按实际标注调整 mask[mask <= 127] = 0.0 if self.transform: augmentations = self.transform(image=image, mask=mask) image = augmentations['image'] mask = augmentations['mask'] return image, mask

这段代码的关键参数是二值化阈值。细胞掩膜如果是 8 位灰度图,通常用 127 作为阈值,但有些标注工具导出的是 0 和 255,有些是 0 和 1,你得先确认。另一个坑是图像和掩膜的文件名对应关系,如果掩膜命名规则和图像不一致,replace那行就会找不到文件。我一般会先写个小脚本统计一下图像和掩膜的数量是否一致,文件名是否能一一对应,再开始训练。

3. 把环境跑起来:依赖安装、训练与评估的完整链路

3.1 环境准备与依赖安装

这份包给了requirements.txt和Dockerfile,说明作者至少考虑过环境问题。我一般优先用 Docker,因为医学图像分割涉及 PyTorch、CUDA、OpenCV、NumPy 等一堆依赖,版本对不上就是各种玄学报错。如果不用 Docker,那就老老实实建虚拟环境。

# 创建虚拟环境,Python 版本建议 3.8 或 3.9 python -m venv venv source venv/bin/activate # Windows 用 venv\Scripts\activate # 安装依赖 pip install -r requirements.txt # 如果 requirements.txt 里没有固定版本,建议手动确认几个关键库 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install opencv-python numpy pillow matplotlib

这里有个血泪经验:requirements.txt里如果写了torch但没有指定 CUDA 版本,pip 默认会装 CPU 版,训练速度慢到你想砸键盘。装完之后一定要验证:

import torch print(torch.__version__) print(torch.cuda.is_available()) # 必须是 True print(torch.cuda.get_device_name(0))

如果cuda.is_available()返回 False,先检查显卡驱动,再检查 CUDA 版本和 PyTorch 版本是否匹配。这一步不通过,后面训练全是白搭。

3.2 训练脚本怎么读、参数怎么改

train.py是训练入口。我一般先不急着跑,而是把文件从头到尾读一遍,重点看这几个地方:数据集路径、模型选择、损失函数、优化器、学习率、epoch 数、batch size、保存路径。

# 常见训练命令,具体参数名以 train.py 的 argparse 为准 python train.py \ --data_dir ./data/cells \ --model unet \ --epochs 50 \ --batch_size 4 \ --learning_rate 1e-4 \ --val_percent 0.1 \ --save_dir ./checkpoints

参数说明:--model通常可以选unet或unet_plus_plus,对应两套模型;--batch_size在显存允许的情况下尽量大一点,细胞图像分割常用 4 到 8;--learning_rate用 1e-4 或 1e-3 起步,如果 loss 震荡就降到 1e-5;--val_percent是验证集比例,0.1 表示 10% 数据用于验证。如果train.py里没有这些参数,而是硬编码在文件里,那就直接改源码,但记得改完备份一份。

损失函数方面,细胞分割常用 Dice Loss 加 BCE Loss 的组合。utils/dice_score.py里应该实现了 Dice 系数计算,训练时可能直接调用。Dice Loss 对类别不平衡不敏感,细胞图像里背景通常远多于细胞,所以 Dice Loss 比纯交叉熵更合适。如果训练时 loss 降到某个值就不动了,可以试试把 BCE 的权重调低,或者加一个 Focal Loss。

3.3 评估与推理:evaluate.py 和 predict.py 怎么用

训练完之后,evaluate.py用来算验证集上的指标,通常是 Dice 系数和 IoU。predict.py用来对单张或一批图像做推理,输出分割掩膜。

# 评估 python evaluate.py \ --model ./checkpoints/best_model.pth \ --data_dir ./data/cells/val \ --model_type unet # 单图推理 python predict.py \ --model ./checkpoints/best_model.pth \ --input ./data/test/cell_001.png \ --output ./results/cell_001_mask.png \ --model_type unet

evaluate.py里一般会加载模型权重,遍历验证集,计算 Dice 和 IoU。如果 Dice 低于 0.7,先别急着调模型,检查一下验证集的掩膜阈值和训练集是否一致。我遇到过训练集掩膜是 0/1,验证集掩膜是 0/255,结果 Dice 只有 0.3,改完阈值直接上 0.85。predict.py的输出通常是二值图,有些实现会输出概率图,需要自己加阈值。如果输出全黑,先看模型是否加载成功,再看输入图像的归一化是否和训练时一致。

3.4 SAHI 切片推理:大图细胞分割的关键一步

细胞图像往往很大,比如 2000x2000 甚至更大,直接缩放到 512x512 送进网络,小细胞就变成几个像素,分割效果很差。slicePredict.py调用 SAHI 做切片推理,把大图切成有重叠的小块,每块单独预测,再拼回原图。

# 切片推理,具体参数以 slicePredict.py 为准 python slicePredict.py \ --model ./checkpoints/best_model.pth \ --input ./data/large/cells_big.png \ --output ./results/cells_big_mask.png \ --slice_size 512 \ --overlap 0.2 \ --model_type unet_plus_plus

--slice_size是每块的大小,通常和训练时的输入尺寸一致;--overlap是块之间的重叠比例,0.2 表示 20% 重叠,重叠是为了避免边界处的细胞被切断。sahi/postprocess.py里会做非极大值抑制或加权融合,把各块的预测拼起来。如果拼接后出现明显的网格状接缝,说明 overlap 太小,调到 0.3 或 0.4 试试。如果显存够,slice_size 可以调到 1024,但要注意训练时用的尺寸,推理尺寸和训练尺寸差太多也会掉点。

4. 避坑与排查:细胞分割训练中最容易翻车的五个地方

4.1 现象:训练 loss 不降,Dice 一直在 0.1 附近

原因:最常见的是掩膜阈值不对。训练集掩膜如果是 0 和 255,但代码里按 0 和 1 处理,模型学到的全是背景。另一个原因是图像和掩膜没有对齐,比如图像做了归一化但掩膜没有,或者数据增强时图像和掩膜用了不同的变换。

解决:先写个脚本可视化几张训练图像和对应的掩膜,确认掩膜里细胞区域是白色,背景是黑色。然后检查data_loading.py里的二值化逻辑,确保掩膜值被正确映射到 0 和 1。如果用了 albumentations 做增强,确认image和mask传的是同一个变换。

4.2 现象:训练正常,推理时输出全黑或全白

原因:推理时的预处理和训练时不一致。比如训练时图像做了ToTensor和Normalize,推理时只做了ToTensor,没做Normalize,模型看到的输入分布变了。或者模型加载时state_dict的 key 不匹配,实际加载的是随机权重。

解决:把推理脚本里的预处理步骤和训练脚本里的逐行对比,确保完全一致。加载模型后打印一下第一层权重的均值,如果接近 0 且方差很小,说明权重没加载成功。另外检查model.eval()和torch.no_grad()是否加了,这两个不加会影响推理结果。

4.3 现象:SAHI 切片推理后,拼接处细胞被切成两半

原因:切片之间的重叠不够,或者后处理融合策略有问题。如果 overlap 设成 0.1,而细胞直径比较大,边界处的细胞就会被切断。另外,如果后处理直接取最大值而不做加权,接缝处会出现硬边。

解决:把 overlap 调到 0.3 到 0.4,让相邻块有足够的重叠区域。然后检查sahi/postprocess.py里的融合逻辑,常见做法是对重叠区域做加权平均,权重从中心到边缘递减。如果代码里没有加权,可以自己加一个高斯权重图。

4.4 现象:显存不够,batch size 只能设 1,训练太慢

原因:UNet++ 参数量比 UNet 大,加上高分辨率输入,显存很容易爆。另外,如果数据加载时没有用pin_memory和num_workers,GPU 利用率也会很低。

解决:先用 UNet 跑,UNet++ 作为对比实验。如果必须用 UNet++,可以减小输入尺寸,比如从 512 降到 256,但要注意小细胞可能丢失。另外,把batch_size设小,用梯度累积模拟大 batch。数据加载方面,num_workers设成 4 或 8,pin_memory=True,能明显提升 GPU 利用率。

4.5 现象:Dice 系数在验证集上很高,但测试集上很差

原因:过拟合。细胞图像数据集通常不大,如果验证集和训练集来自同一批数据,分布相似,验证集 Dice 会虚高。另外,如果数据增强太弱,模型学到的是训练集的特定模式,换一批染色条件不同的图像就崩了。

解决:把数据增强加强,加随机旋转、翻转、弹性变形、颜色抖动。如果条件允许,用不同来源的细胞图像做交叉验证。另外,早停法比固定 epoch 更靠谱,监控验证集 Dice,连续几个 epoch 不提升就停。

5. 进阶技巧:用 UNet++ 深监督和 Dice 阈值搜索把边界抠准

5.1 深监督怎么用,为什么对细胞边界有效

UNet++ 的一个核心特性是深监督,也就是在嵌套解码器的每一层都输出一个预测,然后和真值算 loss,最后把各层 loss 加权求和。这样做的好处是,浅层解码器直接收到梯度,能更好地学习边缘细节。细胞分割最怕的就是边界糊成一团,深监督能让浅层特征也参与到边界恢复中。

在unet_model.py里,UNet++ 类通常会返回一个列表,包含各层的输出。训练时,train.py里会遍历这个列表,对每个输出算 loss。如果你发现边界不够准,可以调高浅层 loss 的权重。常见做法是各层权重相等,或者从深到浅递减。我一般会先跑默认权重,如果边界 Dice 低,再把浅层权重调高 0.2 左右试试。

# 深监督 loss 的常见写法,具体以包内 train.py 为准 outputs = model(inputs) # UNet++ 返回多尺度输出列表 loss = 0 weights = [1.0, 0.8, 0.6, 0.4] # 从深到浅,按实际层数调整 for output, weight in zip(outputs, weights): loss += weight * criterion(output, masks) loss = loss / sum(weights)

这段代码的关键是weights列表的长度要和outputs一致。如果 UNet++ 有 4 层输出,weights就得有 4 个值。权重怎么设没有绝对标准,我的习惯是深层权重高一点,因为深层语义更准,但浅层也不能太低,否则边界恢复不够。

5.2 Dice 阈值搜索:别再用 0.5 一刀切

推理时,模型输出的是概率图,通常用 0.5 作为阈值二值化。但细胞分割里,0.5 不一定最优。如果细胞边界模糊,0.5 可能把边界像素判成背景,导致细胞缩小;如果阈值太低,又会把噪声判成细胞。我一般会在验证集上做阈值搜索,从 0.3 到 0.7,每隔 0.05 试一次,看哪个阈值 Dice 最高。

# 阈值搜索示例 best_dice = 0 best_threshold = 0.5 for threshold in np.arange(0.3, 0.75, 0.05): dice_scores = [] for img, mask in val_loader: pred = model(img) pred_binary = (pred > threshold).float() dice = dice_coeff(pred_binary, mask) dice_scores.append(dice.item()) avg_dice = np.mean(dice_scores) if avg_dice > best_dice: best_dice = avg_dice best_threshold = threshold print(f"Best threshold: {best_threshold}, Best Dice: {best_dice}")

这个搜索过程不复杂,但很实用。我见过不少项目直接用 0.5,Dice 卡在 0.8 上不去,搜完阈值能到 0.85 甚至 0.88。注意,阈值搜索要在验证集上做,不能在测试集上做,否则就是作弊。搜完之后,把最优阈值固定到推理脚本里。

5.3 后处理:去掉小连通域和填洞

细胞分割的输出往往有一些小噪点,或者细胞内部有空洞。sahi/postprocess.py或utils/utils.py里可能有后处理函数,如果没有,可以自己加。常见做法是用连通域分析去掉面积小于某个阈值的区域,再用形态学闭运算填洞。

import cv2 import numpy as np def postprocess_mask(mask, min_area=50): # mask 是二值图,0 和 1 mask = mask.astype(np.uint8) num_labels, labels, stats, _ = cv2.connectedComponentsWithStats(mask, connectivity=8) cleaned = np.zeros_like(mask) for i in range(1, num_labels): if stats[i, cv2.CC_STAT_AREA] >= min_area: cleaned[labels == i] = 1 # 填洞 kernel = np.ones((3, 3), np.uint8) cleaned = cv2.morphologyEx(cleaned, cv2.MORPH_CLOSE, kernel, iterations=2) return cleaned

min_area是最小连通域面积,根据细胞大小调整。如果细胞直径大约 20 像素,面积约 300,min_area可以设 100 到 150。填洞的闭运算核大小也要看细胞间隙,核太大会把相邻细胞粘在一起。我一般会先可视化几张后处理前后的图,确认没有把真细胞滤掉,也没有把背景噪点留下。

5.4 一个我踩过的坑:验证集 Dice 高但实际计数不准

有一次我跑完 UNet++,验证集 Dice 0.89,看着挺美,结果拿去做细胞计数,数量比人工标注多了 30%。排查后发现,模型把一些染色较深的背景区域也判成了细胞,这些区域面积小,Dice 对面积不敏感,所以指标没掉,但计数全乱了。后来我在后处理里加了面积过滤和圆形度过滤,计数才准。

从那以后我每次做完分割,都会强制走一遍计数验证:随机抽几张图,把模型预测的细胞数量和人工标注的数量对比,误差超过 10% 就回去查后处理。Dice 高不代表下游任务准,这是血泪教训。希望帮到你。

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

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

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

立即咨询