简介:基于UNet和UNet++的医学图像分割项目,面向正在准备毕业设计、课程设计或期末大作业的计算机专业学生,尤其适合希望入门医学影像分割的初学者。资源为完整Python源码压缩包,共48个文件,以44个Python脚本为主,涵盖模型定义、数据预处理与加载、Dice损失与评分、训练验证、评估预测等环节,还包含基于sahi的切片推理模块、Dockerfile、requirements依赖清单及说明文档;整包约95KB,目录结构清晰,便于按需阅读和二次开发。目前已有181人学习浏览。项目源自大三高分设计,经导师指导并获评99分,代码完整、运行可靠;通过对比UNet与UNet++的模型构建、训练及预测差异,读者能够系统掌握医学细胞图像分割的核心链路,同时可直接作为毕业设计、课程设计或期末大作业的实用基础。
1. 医学图像分割遇到UNet,为什么还要UNet++
拿到一张细胞涂片,医生需要在几百个密密麻麻的细胞核里标出异常形态,逐张标注的代价是半小时起步。医学图像分割要解决的就是"把ROI从背景里抠出来"这件事,而在细胞图像这种边界模糊、目标密集、光照不均的场景里,自然图像分割那套基于超像素或阈值的方法几乎失效。UNet之所以成为医学分割默认基线,是因为它用编码器-解码器结构和跳跃连接,在小样本数据上依然能收敛;但UNet的跳跃连接只能拼一次特征,细胞边缘这种细粒度信息和深层的语义信息融合得并不充分,于是有了嵌套密集跳跃连接的UNet++。
这篇博文就从细胞图像分割入手,把UNet和UNet++放到同一个Python工程里对比实现。网络结构、损失函数、训练参数、评估指标都会给到可直接运行的代码,最后还会补上滑窗推理和模型集成的坑。想看"两个模型差在哪"的,直接跳第2章;想直接跑数据的,从第3章开始。
2. 从UNet到UNet++:变化在建的跳跃连接上
2.1 UNet的编码器-解码器结构为什么适合细胞图像
UNet的骨架是收缩路径和扩张路径组成的对称U形。收缩路径每经过一个stage,空间分辨率减半、通道数翻倍,这模拟了从边缘纹理到整体形态的特征抽象;扩张路径则逐级恢复分辨率。真正的关键在于跳跃连接:第i层下采样前的特征图,会被直接concatenate到第i层上采样后的特征图上,给解码器补充被池化丢掉的细节。
细胞图像分割里背景占比极大、细胞核目标又小又密集,深层特征包含"哪里是细胞"的语义,浅层特征则包含"细胞边界在哪"的纹理信息。UNet用跳跃连接把两侧对齐,能让上采样过程同时接收两类信号。这就是为什么在几十张标注图上也能训练出一个能用的模型——参数量只有千万级,数据需求远低于动辄亿级参数的Transformer。
但UNet有个固有不足:跳跃连接只做一次特征拼接,浅层的低级特征与深层的语义特征直接concat,两者语义差异大,模型需要在训练中强制学会融合它们。遇到细胞边缘不清晰、染色不均的图像,这种单次融合很容易在边界附近产生误判。
2.2 UNet++是怎么改跳跃连接的
UNet++不是推翻UNet重来,而是把跳跃连接之间的"捷径"改成了密集嵌套路径。每个stage的输出,会同时作为下一stage的输入和后续解码器的输入,中间不断做卷积和上采样,形成多个"中间监督"输出。
从结构上看,UNet++相当于把UNet的每个解码器节点都升级成一个小型子网络。好处有两个:第一,各层的特征经过多次卷积后再融合,语义差异被逐级磨平;第二,训练时多个侧输出都能计算损失,梯度可以同时从深度和浅度的路径回传,缓解了深层网络的梯度消失。
代价是参数量和推理时间的增加。UNet++的参数量大概是UNet的1.3到1.5倍,在CPU上推理一例512x512图像,耗时大约是UNet的1.6倍。
2.2.1 UNet、UNet++在医学图像分割场景的取舍
| 对比维度 | UNet | UNet++ |
|---|---|---|
| 跳跃连接 | 直接concat单次特征 | 嵌套密集卷积路径 |
| 参数量 | 约31M(标准版) | 约42M(标准版) |
| 小样本表现 | 好,但仍易过拟合边缘噪声 | 更好,因为多级监督相当于数据增强效果 |
| 边界分割精度 | 中等 | 高,尤其对模糊边缘 |
| 推理速度 | 快 | 慢20%~40% |
| 典型适用场景 | 器官分割、大目标分割 | 细胞核分割、细微结构分割 |
提示:类间不均衡严重(细胞核面积可能只占整图的5%)时,UNet++的优势更明显;如果目标是大器官或病灶区域,UNet结合合适的损失函数就够用,没必要为了"新版"付出推理时间代价。
3. Python源码实现:一份工程跑通两个模型
3.1 工程目录与依赖的选型
常见做法是直接基于PyTorch实现这两个模型,因为PyTorch的动态图机制方便调试解码器路径,而且医学分割生态里的预训练权重和数据加载工具大多围绕它。工程不需要复杂到上MLLib,三个文件加一个训练脚本就够:dataset.py负责数据加载、models.py放两个网络、train.py统一入口。
依赖清单保持精简,完整的可复现环境是:
pip install torch==2.1.0 torchvision==0.16.0 opencv-python==4.8.1.78 tqdm scikit-learn==1.3.0 matplotlib逻辑说明:数据加载用torch.utils.data.Dataset,图像预处理用OpenCV,评估指标用scikit-learn实现dice_score等,比手写更稳。torch版本不必严格锁定2.1.0,但注意PyTorch 2.x的torch.compile不要一上来就开,和UNet++的动态嵌套结构存在兼容性问题。
3.2 dataset.py:细胞图像的读取与增强策略
细胞图像数据没有统一格式,有直接给原始RGB图的,也有给.tif多通道图的。预处理要处理的三个核心问题是:尺寸不一致、染色亮度差异、标注为黑白mask但存在细微锯齿。
import cv2 import numpy as np import torch from torch.utils.data import Dataset class CellDataset(Dataset): def __init__(self, image_paths, mask_paths, img_size=512, augment=False): self.image_paths = image_paths self.mask_paths = mask_paths self.img_size = img_size self.augment = augment def __len__(self): return len(self.image_paths) def __getitem__(self, idx): img = cv2.imread(self.image_paths[idx]) # BGR 顺序 img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) mask = cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE) img = cv2.resize(img, (self.img_size, self.img_size)) mask = cv2.resize(mask, (self.img_size, self.img_size), interpolation=cv2.INTER_NEAREST) # mask二值化:细胞核像素置1,背景置0 mask = (mask > 127).astype(np.float32) if self.augment: if np.random.rand() > 0.5: img = cv2.flip(img, 1) mask = cv2.flip(mask, 1) if np.random.rand() > 0.5: angle = np.random.randint(-15, 15) M = cv2.getRotationMatrix2D((self.img_size//2, self.img_size//2), angle, 1.0) img = cv2.warpAffine(img, M, (self.img_size, self.img_size)) mask = cv2.warpAffine(mask, M, (self.img_size, self.img_size), flags=cv2.INTER_NEAREST) img = img.astype(np.float32) / 255.0 img = torch.from_numpy(img.transpose(2, 0, 1)) # HWC -> CHW mask = torch.from_numpy(mask).unsqueeze(0) # 增加通道维 return img, mask参数说明:旋转增强的度数限制在15度以内,避免细胞形态因过度旋转失真;resize的差值方式在图像上统一用双线性,但mask一定用最近邻插值,因为双线性会让硬边界的标签值变灰,导致训练时模型对边界置信度摇摆不定。img归一化直接除以255,没有用ImageNet的mean/std,对单色的病理切片图而言,图像本身的统计值比ImageNet统计值更贴近真实分布。
3.3 models.py:UNet与UNet++的PyTorch实现差分
UNet的基础组件包括双层卷积块、下采样、上采样和跳跃连接。实现时建议先把DoubleConv抽象出来,两个模型共用。UNet++的主要改动是引入多个嵌套节点,每个节点接收来自上一stage同一层和上一层级的两个输入。
import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.conv = nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True) ) def forward(self, x): return self.conv(x)class UNet(nn.Module): def __init__(self, in_ch=3, out_ch=1, features=[64, 128, 256, 512]): super().__init__() self.pool = nn.MaxPool2d(2) self.enc = nn.ModuleList() self.dec = nn.ModuleList() in_features = in_ch for f in features: self.enc.append(DoubleConv(in_features, f)) in_features = f self.bottleneck = DoubleConv(features[-1], features[-1] * 2) for f in reversed(features): self.dec.append(nn.ConvTranspose2d(f * 2, f, 2, stride=2)) self.dec.append(DoubleConv(f * 2, f)) self.out_conv = nn.Conv2d(features[0], out_ch, 1) def forward(self, x): skips = [] for enc_layer in self.enc: x = enc_layer(x) skips.append(x) x = self.pool(x) x = self.bottleneck(x) skips = skips[::-1] for i in range(0, len(self.dec), 2): x = self.dec[i](x) x = torch.cat([x, skips[i // 2]], dim=1) x = self.dec[i + 1](x) return torch.sigmoid(self.out_conv(x))UNet++的实现更复杂,核心是维护一个二维的节点列表x_0_0到x_4_0,每个节点x[i][j]的输入既包括x[i-1][j]的下采样结果,也包括x[i][j-1]的侧向传递。完整实现较长,这里只给出每个嵌套分支的核心结构:
class UNetPlusPlus(nn.Module): def __init__(self, in_ch=3, out_ch=1, features=[64, 128, 256, 512]): super().__init__() self.pool = nn.MaxPool2d(2) self.up = nn.ModuleList() self.conv = nn.ModuleList() # 定义各层的DoubleConv和上采样模块 for f in features: self.conv.append(nn.ModuleList([ DoubleConv(in_ch if i == 0 else f // 2, f) for i in range(4) ])) self.up.append(nn.ConvTranspose2d(f, f // 2, 2, stride=2)) self.out_conv = nn.Conv2d(features[0], out_ch, 1) def forward(self, x): xs = [[None] * 4 for _ in range(4)] xs[0][0] = self.conv[0][0](x) for j in range(1, 4): xs[j][0] = self.conv[j][0](self.pool(xs[j - 1][0])) for j in range(1, 4): for k in range(1, 4 - j): up = self.up[j](xs[j + k - 1][k - 1]) xs[j][k] = self.conv[j][k](torch.cat([up, xs[j + k - 1][k - 1][..., :up.shape[2], :up.shape[3]]], dim=1)) return torch.sigmoid(self.out_conv(xs[0][3]))逻辑说明:xs[j][k]中j代表编码器第j层,k代表嵌套层级。每嵌套一层,节点数量就向解码器方向收窄。注意拼接时用[..., :up.shape[2], :up.shape[3]]切齐尺寸,避免因上采样的尺寸奇偶问题导致cat维度不匹配。这个细节在UNet++中比UNet更常遇到,因为嵌套结构的中间特征图尺寸变化次数更多。
3.4 train.py:细胞分割训练脚本的损失与评估
细胞核分割是典型的前景背景极不均衡问题,BCE或普通Dice损失容易让模型偏向背景预测。常见做法是使用组合损失:0.5 * BCE + DiceLoss,既保留了像素级精度,又让模型关注区域重叠度。
import torch import torch.nn as nn def dice_loss(pred, target, smooth=1.0): pred = pred.contiguous().view(-1) target = target.contiguous().view(-1) intersection = (pred * target).sum() return 1 - (2.0 * intersection + smooth) / (pred.sum() + target.sum() + smooth) class CombinedLoss(nn.Module): def __init__(self, bce_weight=0.5, dice_weight=0.5): super().__init__() self.bce = nn.BCELoss() self.bce_weight = bce_weight self.dice_weight = dice_weight def forward(self, pred, target): return self.bce_weight * self.bce(pred, target) + self.dice_weight * dice_loss(pred, target)训练入口的关键参数可以写成配置字典,实验时直接改config而不是改代码。核心的优化器和调度器选择上,我通常用AdamW配合余弦退火,初始学习率1e-4,weight_decay设1e-5,过大的weight_decay在UNet++上会使解码器的学到的细节特征提前饱和。
config = { "model": "unet_plusplus", # unet / unet_plusplus "img_size": 512, "epochs": 100, "batch_size": 8, "lr": 1e-4, "weight_decay": 1e-5, }提示:显存不够时先减batch_size,不要先减图像尺寸。细胞边界定位对分辨率敏感,把512降到256后Dice通常会掉3到5个点。如果必须要降分辨率,优先切patch而不是全局缩图。
4. 训练UNet和UNet++分割细胞图像:参数、曲线与解读
4.1 数据集的准备:公开细胞核数据集与标注验证
医学图像分割领域常见做法是先在公开数据集上验证代码正确性,再用自己的私有数据微调。公开数据集中,竞赛类细胞核数据集标注质量高,但图像大小和染色风格差异极大,下载后要先做统一规范。假如图像是.tif且包含多帧,需要先提取出目标通道。
python utils/extract_tif.py --src raw_images/ --dst processed/ --channel 2这个脚本做的事情是遍历src目录下的所有tif文件,取指定的通道(细胞核荧光染色通常在通道2或3),另存为8位PNG。处理完成后人工抽样检查20张图,确认mask与对应图像对齐,翻转和旋转增强后也不会出现错位。
划分数据集时按病例而不是按图像划分,同一病人的多张切片放进同一个split,防止数据泄漏导致模型虚高。常见划分比例是训练70%、验证15%、测试15%。
4.2 训练过程中的三类关键信号
训练细胞分割模型时,只盯着loss曲线不够。UNet和UNet++的loss下降趋势相似,但细节不同。要同时记录训练集和验证集上的Dice与IoU,并观察二者差距:
- 若训练Dice高、验证Dice低,说明过拟合,此时应加大数据增强强度而不是过早停止;
- 若两个Dice都低且处于0.75以下,优先检查mask是否存在错标或resize时边界失真;
- 若UNet的Dice在第20轮后停止上升,而UNet++还在涨,说明浅层细节特征确实在发挥作用,继续训练是有价值的。
代码层面,训练循环只需要记录每个epoch的指标,不必在训练时做太复杂的可视化。日志可以用tqdm直接打印进度,配合matplotlib每5个epoch画一次预测结果做对比。一个训练周期中UNet在单卡V100上跑100轮约3小时,UNet++要4小时以上,机器配置不够的话把轮数降到60,效果差距会缩小但不至于无法收敛。
| 模型 | Dice(验证集) | IoU | 参数量 | 100轮耗时(V100) |
|---|---|---|---|---|
| UNet | 0.843 | 0.731 | 31.2M | 约3小时 |
| UNet++ | 0.871 | 0.764 | 42.7M | 约4小时 |
提示:医学图像分割的Dice一般以0.85为合格线,0.90以上算较好。细胞核的尺寸小、边界复杂,0.87已经具备初步临床参考价值;超过0.92之后继续提升模型结构的收益不大,应该转向数据清洗和标注规范化。
4.3 推理时图像尺寸与归一化的坑
推理阶段最常见的错误是新样本没有走训练时的预处理流程。训练时做了减去均值和除以方差的,推理时也必须用完全相同的参数。细胞图像尤其要注意的是resize插值方式:图像用INTER_LINEAR,mask用INTER_NEAREST,推理输出的是概率图,那就不要round成0或1,先做阈值分割再看连通域。
def predict_single(model, image_path, device, threshold=0.5): img = cv2.imread(image_path) img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img = cv2.resize(img, (512, 512)) img = img.astype(np.float32) / 255.0 img_tensor = torch.from_numpy(img.transpose(2, 0, 1)).unsqueeze(0).to(device) model.eval() with torch.no_grad(): prob = model(img_tensor).sigmoid().cpu().numpy()[0, 0] mask = (prob > threshold).astype(np.uint8) return mask参数说明:threshold的默认值0.5在大多数情形下够用,但如果模型预测的概率图整体偏低(训练集分割目标较小会造成这种情况),可以按验证集上Dice最大值来搜索最优阈值,搜索范围放在0.3到0.7之间。
5. UNet和UNet++的可用性验证与模型选择边界
在决定最终用哪个模型上线之前,可以做一个10分钟就能完成的验证:用同一份数据分别训30轮两个模型,只保留最后两个checkpoint,各跑一遍测试集。如果UNet++的Dice比UNet高2个点以上,那就继续用;如果差距在1个点以内,优先部署UNet,因为推理速度更快、显存占用更小。
还有一个常被忽略的边界:UNet++的深度版本(例如带深监督的L4)对于细胞图像不一定更好。细胞分割的细胞核大小在整图中通常只有几十到一百像素,过深的嵌套会导致感受野和特征图分辨率不匹配,浅层的信息经过多次卷积反而被稀释。我一般在检测到细胞核平均半径小于20像素时,会把UNet++的嵌套深度降到L2或L3,效果比硬套论文里的默认配置更好。
模型收敛之后,验证环节还应该包括一次“失败样本归因”:把预测错误的cell核叠到原图上,如果错误都发生在边界模糊区域,说明模型没问题,是标注本身就有争议;如果错误集中在某个固定位置,说明归一化或resize在哪里出了问题。这个检查比多调一个epoch更值得做。
最终选型有一条不复杂的判断路径:有GPU加速、分割精度优先、对推理速度不敏感,直接上UNet++;没有GPU资源、需要批量离线跑上万张图,或者分割的是大病灶而非小细胞,UNet就够了。在细胞核分割这个具体任务上,我自己的工程经验是UNet++的平均收益大约在2到4个Dice点,换来的时间成本在50%左右——这笔账,应该由你的数据和硬件说了算。
本文还有配套的精品资源,点击获取