简介:这份资源面向计算机相关专业正在做毕业设计、课程设计或期末大作业的学生,以及需要医学图像分割实战练习的学习者,提供基于UNet与UNet++两种网络结构对细胞图像进行分割的完整Python实现,帮助读者理解语义分割在医学影像场景中的落地流程。压缩包共48个文件,以44个py源码为主,另含requirements.txt依赖清单、Dockerfile容器配置、readme.md说明文档及.gitignore等辅助文件,整体约95KB,体积轻便便于本地部署与二次修改。代码涵盖数据加载与预处理、Dice系数评估、模型定义、训练与预测脚本,并集成切片推理与后处理模块,目录按unet、sahi、scripts、utils等分层组织,结构清晰。目前已有181人学习下载,适合作为分割项目入门与工程化参考,帮助读者快速跑通训练与推理链路,掌握从数据到评估的完整实现思路。
1. 细胞图像分割为什么总在边界上翻车:UNet 与 UNet++ 要解决的真实问题
做过细胞图像分割的人大多有过类似经历:模型在训练集上 Dice 看着不错,一放到新批次染色切片上,细胞之间的粘连区域就糊成一片,边界像被橡皮擦抹过。这不是调参不够勤快,而是细胞图像本身的特性决定的——目标密集、边界对比度低、不同染色批次分布漂移。UNet 和 UNet++ 这一对编码器-解码器结构,正是为这种「少样本、强边界、多尺度」场景设计的经典方案。UNet 用跳跃连接把浅层高分辨率特征直接送到解码端,缓解下采样造成的边界信息丢失;UNet++ 则在跳跃连接上再嵌套密集卷积块和深监督,让不同语义尺度的特征在解码前先对齐。这套 Python 源码要落地的,就是一套能直接跑自己细胞数据集、能对比两种结构差异、能导出可视化掩膜的训练流程。适合已经会写 PyTorch 训练循环、但被细胞粘连和边界模糊卡住的从业者,也适合想拿医学图像分割当第一个语义分割实战项目的新手。
2. 把 UNet 和 UNet++ 拆开看:结构差异与选型理由
2.1 UNet 的跳跃连接到底补了什么
UNet 的结构可以粗暴理解成「下采样五次、上采样五次、每次上采样前把对应下采样层的特征拼过来」。下采样负责扩大感受野、提取语义,上采样负责恢复分辨率,而跳跃连接负责把下采样过程中被池化丢掉的边缘、纹理信息重新注入。细胞图像里细胞核边界往往只有几个像素宽,如果只靠深层语义特征上采样,边界必然模糊。跳跃连接让解码器在恢复尺寸时能直接看到原始分辨率下的梯度变化,这是 UNet 在细胞分割上长期作为基线的核心原因。
但 UNet 的跳跃连接是「硬拼接」:编码器第 i 层特征直接 concat 到解码器对应层。浅层特征语义弱、噪声多,深层特征语义强、位置粗,两者直接拼接会存在语义鸿沟。细胞图像里表现为:大细胞轮廓还行,小细胞和粘连处容易漏。这不是 UNet 错了,而是它的设计目标本来就是通用分割,没有专门处理多尺度语义对齐。
2.2 UNet++ 用嵌套密集连接和深监督填语义鸿沟
UNet++ 的思路是在编码器和解码器之间架一座「密集连接的桥」。它把原来一条跳跃连接拆成多个节点,每个节点接收同一尺度编码器特征和更浅层解码器节点的输出,逐级融合。这样解码器在每一层拿到的特征,已经过多次跨尺度混合,语义鸿沟被缩小。同时 UNet++ 在多个解码节点上加了深监督,也就是中间层也接损失函数,让浅层解码器被迫学到有判别力的特征,而不是只靠最后一层。
对细胞图像来说,这个改动的直接收益是:粘连细胞的分割边界更贴合,小目标召回率通常比 UNet 高。代价是显存和训练时间增加,因为密集连接让中间特征图数量变多。选型上我的习惯是:数据量小于 500 张、细胞边界要求高、显存够 8GB 以上,优先 UNet++;如果只是快速验证流程或部署端算力紧张,UNet 更稳。
2.3 两种结构的参数量与显存对比
| 结构 | 参数量相对量级 | 输入 256×256 单卡显存 | 训练速度 | 边界表现 |
|---|---|---|---|---|
| UNet | 基准 1x | 约 2.5GB | 快 | 中等,粘连处易糊 |
| UNet++ | 约 1.6~2x | 约 4GB | 慢 30%~50% | 较好,小目标召回高 |
提示:显存数值随 batch size 和 backbone 宽度变化,上表按 batch size 4、基础通道 64 估算,实际以自己环境为准。
2.4 数据准备与目录组织的最小步骤
细胞图像分割常见格式是原图加对应掩膜,掩膜里细胞区域为 1、背景为 0。目录建议按下面组织,训练脚本只认这个结构,换数据集时不用改代码。
dataset/ ├── train/ │ ├── images/ # 细胞原图,png 或 tif │ └── masks/ # 对应二值掩膜,文件名与 images 一致 ├── val/ │ ├── images/ │ └── masks/ └── test/ ├── images/ └── masks/文件名必须一一对应,掩膜建议存成单通道 8 位 png,像素值 0 或 255。如果原始掩膜是 0/1,读入后要归一化到 0/1 再算损失,否则 Dice 计算会出错。常见做法是写一个 Dataset 类,读图时同步做随机翻转、旋转、弹性形变,细胞图像弹性形变增强对边界泛化帮助明显。
3. 用 Python 跑通 UNet 训练:从 Dataset 到 Dice 监控
3.1 写一个能同时喂 UNet 和 UNet++ 的 Dataset
import os import cv2 import numpy as np import torch from torch.utils.data import Dataset class CellSegDataset(Dataset): def __init__(self, root, img_size=256, augment=False): self.img_dir = os.path.join(root, "images") self.mask_dir = os.path.join(root, "masks") self.names = sorted(os.listdir(self.img_dir)) self.img_size = img_size self.augment = augment def __len__(self): return len(self.names) def __getitem__(self, idx): name = self.names[idx] img = cv2.imread(os.path.join(self.img_dir, name), cv2.IMREAD_COLOR) img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) mask = cv2.imread(os.path.join(self.mask_dir, name), 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) # 掩膜二值化并归一化到 0/1,避免 Dice 计算出错 mask = (mask > 127).astype(np.float32) if self.augment: if np.random.rand() < 0.5: img = np.fliplr(img).copy() mask = np.fliplr(mask).copy() if np.random.rand() < 0.5: img = np.flipud(img).copy() mask = np.flipud(mask).copy() img = img.astype(np.float32) / 255.0 img = np.transpose(img, (2, 0, 1)) mask = np.expand_dims(mask, axis=0) return torch.from_numpy(img), torch.from_numpy(mask)这个 Dataset 的关键点有三个:掩膜用最近邻插值缩放,避免边界被插值成灰色;掩膜二值化阈值取 127,兼容 0/255 存储;增强只做翻转,因为细胞图像旋转会改变方向语义,弹性形变建议用 albumentations 单独加。参数img_size控制输入尺寸,细胞图像常用 256 或 512,显存不够就降到 256。
3.2 UNet 最小实现与通道数设置
import torch.nn as nn import torch.nn.functional as F class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.net = 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.net(x) class UNet(nn.Module): def __init__(self, in_ch=3, out_ch=1, base=64): super().__init__() self.d1 = DoubleConv(in_ch, base) self.d2 = DoubleConv(base, base * 2) self.d3 = DoubleConv(base * 2, base * 4) self.d4 = DoubleConv(base * 4, base * 8) self.bottleneck = DoubleConv(base * 8, base * 16) self.pool = nn.MaxPool2d(2) self.up4 = nn.ConvTranspose2d(base * 16, base * 8, 2, stride=2) self.u4 = DoubleConv(base * 16, base * 8) self.up3 = nn.ConvTranspose2d(base * 8, base * 4, 2, stride=2) self.u3 = DoubleConv(base * 8, base * 4) self.up2 = nn.ConvTranspose2d(base * 4, base * 2, 2, stride=2) self.u2 = DoubleConv(base * 4, base * 2) self.up1 = nn.ConvTranspose2d(base * 2, base, 2, stride=2) self.u1 = DoubleConv(base * 2, base) self.out = nn.Conv2d(base, out_ch, 1) def forward(self, x): c1 = self.d1(x) c2 = self.d2(self.pool(c1)) c3 = self.d3(self.pool(c2)) c4 = self.d4(self.pool(c3)) bn = self.bottleneck(self.pool(c4)) x = self.u4(torch.cat([self.up4(bn), c4], dim=1)) x = self.u3(torch.cat([self.up3(x), c3], dim=1)) x = self.u2(torch.cat([self.up2(x), c2], dim=1)) x = self.u1(torch.cat([self.up1(x), c1], dim=1)) return self.out(x)base=64是基础通道数,显存不够改成 32,分割精度会略降但能跑起来。输出层不加 sigmoid,因为损失函数用带 logits 的 BCE,数值更稳。如果细胞图像是灰度图,in_ch改成 1。
3.3 UNet++ 的嵌套解码节点怎么接
class UNetPlusPlus(nn.Module): def __init__(self, in_ch=3, out_ch=1, base=32): super().__init__() # 编码器 self.conv0_0 = DoubleConv(in_ch, base) self.conv1_0 = DoubleConv(base, base * 2) self.conv2_0 = DoubleConv(base * 2, base * 4) self.conv3_0 = DoubleConv(base * 4, base * 8) self.pool = nn.MaxPool2d(2) # 解码节点,每个节点融合同尺度编码和浅层解码 self.conv0_1 = DoubleConv(base * 3, base) self.conv1_1 = DoubleConv(base * 6, base * 2) self.conv2_1 = DoubleConv(base * 12, base * 4) self.conv0_2 = DoubleConv(base * 4, base) self.conv1_2 = DoubleConv(base * 8, base * 2) self.conv0_3 = DoubleConv(base * 5, base) self.up = nn.Upsample(scale_factor=2, mode="bilinear", align_corners=True) self.out = nn.Conv2d(base, out_ch, 1) def forward(self, x): x0_0 = self.conv0_0(x) x1_0 = self.conv1_0(self.pool(x0_0)) x2_0 = self.conv2_0(self.pool(x1_0)) x3_0 = self.conv3_0(self.pool(x2_0)) x0_1 = self.conv0_1(torch.cat([x0_0, self.up(x1_0)], 1)) x1_1 = self.conv1_1(torch.cat([x1_0, self.up(x2_0)], 1)) x2_1 = self.conv2_1(torch.cat([x2_0, self.up(x3_0)], 1)) x0_2 = self.conv0_2(torch.cat([x0_0, x0_1, self.up(x1_1)], 1)) x1_2 = self.conv1_2(torch.cat([x1_0, x1_1, self.up(x2_1)], 1)) x0_3 = self.conv0_3(torch.cat([x0_0, x0_1, x0_2, self.up(x1_2)], 1)) return self.out(x0_3)这里base=32是因为 UNet++ 特征图数量多,base 取 64 容易爆显存。每个convX_Y的输入通道数等于拼接后通道总和,改 base 时这些数字要同步改,否则会报通道不匹配。深监督可以在x0_1、x0_2、x0_3后各接一个 1×1 卷积输出辅助预测,训练时加权求和,推理只用x0_3。
3.4 训练循环与 Dice 监控
import torch from torch.utils.data import DataLoader def dice_loss(logits, target, eps=1e-6): prob = torch.sigmoid(logits) inter = (prob * target).sum(dim=(2, 3)) union = prob.sum(dim=(2, 3)) + target.sum(dim=(2, 3)) return 1 - ((2 * inter + eps) / (union + eps)).mean() train_ds = CellSegDataset("dataset/train", img_size=256, augment=True) val_ds = CellSegDataset("dataset/val", img_size=256, augment=False) train_loader = DataLoader(train_ds, batch_size=4, shuffle=True, num_workers=2) val_loader = DataLoader(val_ds, batch_size=1, shuffle=False) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = UNet(in_ch=3, out_ch=1, base=64).to(device) optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) bce = torch.nn.BCEWithLogitsLoss() for epoch in range(80): model.train() for img, mask in train_loader: img, mask = img.to(device), mask.to(device) logits = model(img) loss = bce(logits, mask) + dice_loss(logits, mask) optimizer.zero_grad() loss.backward() optimizer.step() model.eval() dices = [] with torch.no_grad(): for img, mask in val_loader: img, mask = img.to(device), mask.to(device) prob = torch.sigmoid(model(img)) pred = (prob > 0.5).float() inter = (pred * mask).sum().item() dices.append((2 * inter + 1e-6) / (pred.sum().item() + mask.sum().item() + 1e-6)) print(f"epoch {epoch}, val dice {sum(dices)/len(dices):.4f}")损失用 BCE 加 Dice 组合,BCE 稳定像素级梯度,Dice 直接优化重叠度,细胞图像前景占比低时比单用 BCE 收敛快。学习率 1e-3 配 Adam 是常见起点,验证 Dice 连续 10 轮不升就降到 1e-4。batch_size=4是 8GB 显存下 256 输入的保守值,显存够可以加到 8。
4. 细胞分割避坑排查:从黑边到 Dice 虚高的 5 个血泪记录
4.1 现象:验证 Dice 很高,测试集一塌糊涂
原因:训练集和验证集来自同一批次染色,分布几乎一致,模型记住了染色风格而不是细胞结构。解决:按批次划分 train/val/test,或者用不同来源的细胞数据做外部验证。我一般会留一个完全独立来源的小测试集,哪怕只有 20 张,也能提前暴露泛化问题。
4.2 现象:掩膜边缘出现一圈黑边,Dice 卡在 0.7 上不去
原因:掩膜缩放用了双线性插值,边界被插成 0 到 255 之间的灰值,二值化后边界偏移。解决:掩膜 resize 必须用INTER_NEAREST,读入后先二值化再归一化。这个坑在细胞图像里特别明显,因为细胞边界本来就窄,插值误差直接吃掉一两个像素。
4.3 现象:UNet++ 训练 loss 震荡,显存偶尔爆
原因:密集连接让中间特征图通道数随 base 平方级增长,base=64 时某些节点输入通道超过 1024。解决:UNet++ 的 base 从 32 起步,或者用分组卷积压缩中间通道。如果还是爆,把输入从 512 降到 256,细胞图像 256 通常够用。
4.4 现象:预测结果全是背景或全是前景
原因:细胞图像前景占比低时,BCE 被背景像素主导,模型倾向全预测背景;如果掩膜归一化没做,目标值变成 0/255,sigmoid 输出永远追不上。解决:确认掩膜是 0/1,损失用 BCE 加 Dice,Dice 对类别不平衡不敏感。另外可以在 Dataset 里统计前景占比,低于 5% 时考虑加权采样。
4.5 现象:推理时单张图正常,批量推理结果错位
原因:批量推理时没有对每张图单独做 resize 和归一化,或者用了batch_size>1但模型里有 BatchNorm 在 eval 模式下统计量不对。解决:推理统一batch_size=1,或者确认model.eval()已调用。细胞图像尺寸不一致时,逐张 resize 再拼 batch,不要直接 stack 原图。
5. 把 UNet++ 深监督用起来:一个提升小细胞召回的具体技巧
深监督是 UNet++ 里最容易被忽略、但对细胞图像最实用的部分。细胞图像里小细胞和粘连细胞往往只占几十个像素,最后一层解码特征经过多次上采样后,这些小目标的响应已经被稀释。深监督让中间解码节点也接损失,等于强迫网络在浅层就学会区分小细胞边界。具体做法是在x0_1、x0_2、x0_3后各加一个 1×1 卷积输出辅助 logits,训练时把主输出和辅助输出的损失加权求和,推理时只取主输出。
class UNetPlusPlusDeepSup(UNetPlusPlus): def __init__(self, in_ch=3, out_ch=1, base=32): super().__init__(in_ch, out_ch, base) self.aux1 = nn.Conv2d(base, out_ch, 1) self.aux2 = nn.Conv2d(base, out_ch, 1) self.aux3 = nn.Conv2d(base, out_ch, 1) def forward(self, x): x0_0 = self.conv0_0(x) x1_0 = self.conv1_0(self.pool(x0_0)) x2_0 = self.conv2_0(self.pool(x1_0)) x3_0 = self.conv3_0(self.pool(x2_0)) x0_1 = self.conv0_1(torch.cat([x0_0, self.up(x1_0)], 1)) x1_1 = self.conv1_1(torch.cat([x1_0, self.up(x2_0)], 1)) x2_1 = self.conv2_1(torch.cat([x2_0, self.up(x3_0)], 1)) x0_2 = self.conv0_2(torch.cat([x0_0, x0_1, self.up(x1_1)], 1)) x1_2 = self.conv1_2(torch.cat([x1_0, x1_1, self.up(x2_1)], 1)) x0_3 = self.conv0_3(torch.cat([x0_0, x0_1, x0_2, self.up(x1_2)], 1)) if self.training: return self.out(x0_3), self.aux1(x0_1), self.aux2(x0_2), self.aux3(x0_3) return self.out(x0_3)训练时辅助损失权重建议 0.3、0.2、0.1,主损失权重 1.0。权重太大浅层会主导,边界反而变糙;太小等于没加。这个技巧我在细胞核分割上试过,小目标召回能提 3 到 5 个百分点,代价是训练显存多约 15%。验证时只看主输出 Dice,辅助输出只参与训练。
另一个实用习惯是保存验证 Dice 最高的权重,而不是最后一轮。细胞图像分割的验证曲线经常在中后期震荡,最后一轮未必最好。我一般每轮存一个best.pth,训练结束直接拿它做测试集推理,省得回头翻日志找后悔药。这套 UNet 和 UNet++ 的 Python 流程跑通后,换数据集基本只改 Dataset 路径和in_ch,结构代码不用动。希望帮到你。
本文还有配套的精品资源,点击获取