ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

基于UNet与UNet++的细胞医学图像分割Python实现源码

基于UNet与UNet++的细胞医学图像分割Python实现源码 简介这份源码面向计算机相关专业的毕业设计、课程设计及期末综合作业需求提供基于UNet与UNet两种编码器-解码器架构的医学细胞图像分割完整实现采用Python编写原为本科三年级课程设计在导师指导下获99分评价代码结构完整且验证可运行适合不同基础的学习者参考。资源包共58个文件以44个py源码文件为核心覆盖数据加载、数据增强、模型构建、训练与预测评估等模块另含zbak备份文件、Dockerfile、requirements.txt、readme.md及gitignore等配置说明压缩包约107KB目录组织清晰便于按模块查阅。项目详细实现了损失函数配置、dice_score等评估指标与可复现实验环境说明并通过对比UNet与UNet在细胞分割任务中的表现展示不同网络架构在医学图像处理中的特性与优势所有功能模块均配有注释便于理解算法原理与实现细节。目前已有60人学习适合需要项目实践训练或算法复现的开发者。1. 从一张病理切片说起UNet 与 UNet 在细胞分割里到底解决了什么病理科医生在显微镜下数细胞一张 2048×2048 的切片里可能有上千个细胞边界互相粘连细胞核深浅不一。手工勾画一份标注要几十分钟还带主观差异。基于 UNet 与 UNet 的细胞医学图像分割 Python 实现源码要解决的就是把这份重复劳动交给模型输入一张显微图像输出每个像素属于细胞核还是背景的概率图再二值化成掩膜。这套方案适合三类人刚学完 Python 基础语法、想找一个能跑通的图像分割项目练手的人手里有自己显微镜数据、想训练自己数据集的研究生和工程师以及需要快速搭一个分割基线、再谈 unet 模型改进的人。它不挑硬件一张 8GB 显存的消费级显卡就能跑 512×512 的输入。下面从网络结构、数据管线、训练、推理到踩坑按能复现的顺序讲清楚。2. UNet 与 UNet 的结构差异为什么后者在细胞边界上更稳2.1 编码器-解码器与跳跃连接的基本盘UNet 的核心是 U 形结构左侧编码器逐级下采样通道数翻倍、空间尺寸减半把纹理信息压缩成语义特征右侧解码器逐级上采样把低分辨率特征还原回原图尺寸。真正让它work的是同层之间的跳跃连接把编码器的高分辨率细节直接拼到解码器对应层弥补下采样丢掉的边界信息。细胞分割的难点恰好在这里细胞核直径可能只有十几个像素经过三四次下采样后边界信息几乎被抹平。没有跳跃连接解码器只能靠语义猜边界结果就是掩膜糊成一团。UNet 用 concat 拼接同尺度特征等于给解码器留了一条细节通道这是它在医学图像上长期作为基线的根本原因。代价是参数量和显存。标准 UNet 在 512×512 输入下第一层就是 64 通道浅层特征图很大显存占用主要卡在编码器前两层。我一般会把 base_channels 从 64 降到 32精度掉不到一个点显存能省近一半这是新手最容易忽略的调参点。2.2 UNet 的嵌套密集跳跃连接UNet 没有推翻 UNet而是在跳跃连接上做文章。它把原来一条直连的跳跃连接换成一组嵌套的密集卷积块编码器第 i 层到解码器第 j 层之间插入中间节点每个节点融合同层前一节点和下层上采样结果。用一句话概括UNet 的跳跃连接是「直连」UNet 是「带中间处理的密集连接」。这样做的收益是缩小编码器和解码器特征之间的语义鸿沟。UNet 里编码器浅层特征偏纹理、解码器深层特征偏语义直接 concat 会有语义不一致UNet 的中间节点相当于做了几次过渡卷积让拼接的两侧语义更接近。在细胞边界这种细结构上UNet 的 Dice 通常比 UNet 高 1 到 3 个点代价是训练更慢、显存更高。选型建议很直接数据量小、边界要求高、显存够用 UNet追求训练速度、要快速出基线、显存紧张用 UNet。两者代码可以共用同一套数据管线和训练循环只换模型定义这也是我把它们放在同一个源码工程里的原因。2.3 用 PyTorch 定义两个模型的最小代码下面这段是工程里模型定义的核心UNet 和 UNet 共用 DoubleConv 和 Down/Up 模块区别只在跳跃连接的组织方式。import torch import torch.nn as nn class DoubleConv(nn.Module): 两次 3x3 卷积 BN ReLU是 UNet 系列的基本单元 def __init__(self, in_ch, out_ch): super().__init__() self.net nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.net(x) class UNet(nn.Module): def __init__(self, in_ch3, out_ch1, base32): super().__init__() # 编码器4 次下采样通道 base - base*8 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.pool nn.MaxPool2d(2) self.bottleneck DoubleConv(base * 8, base * 16) # 解码器上采样后与同层编码特征 concat self.up4 nn.ConvTranspose2d(base * 16, base * 8, 2, stride2) self.c4 DoubleConv(base * 16, base * 8) self.up3 nn.ConvTranspose2d(base * 8, base * 4, 2, stride2) self.c3 DoubleConv(base * 8, base * 4) self.up2 nn.ConvTranspose2d(base * 4, base * 2, 2, stride2) self.c2 DoubleConv(base * 4, base * 2) self.up1 nn.ConvTranspose2d(base * 2, base, 2, stride2) self.c1 DoubleConv(base * 2, base) self.head nn.Conv2d(base, out_ch, 1) def forward(self, x): e1 self.d1(x) e2 self.d2(self.pool(e1)) e3 self.d3(self.pool(e2)) e4 self.d4(self.pool(e3)) b self.bottleneck(self.pool(e4)) d4 self.c4(torch.cat([self.up4(b), e4], dim1)) d3 self.c3(torch.cat([self.up3(d4), e3], dim1)) d2 self.c2(torch.cat([self.up2(d3), e2], dim1)) d1 self.c1(torch.cat([self.up1(d2), e1], dim1)) return self.head(d1)逻辑说明DoubleConv 是重复两次的卷积块BN 放在卷积后、激活前能稳定训练。编码器每经过一次 pool 空间减半解码器用转置卷积上采样后与编码器同层特征在通道维 concat再过一个 DoubleConv 融合。最后 1×1 卷积把通道压到 out_ch二分类就是 1。参数说明base 控制模型宽度默认 32显存紧张可降到 16追求精度可升到 64in_ch 按输入通道设RGB 是 3灰度显微图是 1out_ch 二分类为 1多类细胞分割改成类别数。UNet 在此基础上把跳跃连接换成嵌套节点工程里通常用一个嵌套层数参数控制深度层数越多显存越高一般取 2 到 3 层就够。3. 数据管线从显微镜图像到可训练张量3.1 细胞分割数据集的目录约定与划分工程里我固定一套目录约定避免路径写死在代码里data/ train/ images/ *.png masks/ *.png val/ images/ masks/ test/ images/images 和 masks 文件名必须一一对应mask 用 0 表示背景、255 表示细胞核这是医学分割最常见的标注格式。划分比例按 7:1.5:1.5 走如果样本量少于 200 张验证集至少留 20 张否则指标抖动大到没法判断模型好坏。常见做法是先按患者或切片划分而不是随机划分图像。同一张切片切出来的图块高度相似随机划分会让验证集泄漏训练信息指标虚高上线就翻车。这一点在细胞分割里比自然图像更严重。3.2 数据增强与归一化细胞图像不能照搬自然图像那套细胞显微图像和自然图像分布差别很大增强策略要克制。水平翻转、垂直翻转、90 度旋转是安全的因为细胞没有固定朝向随机裁剪到 256 或 512 也常用。但颜色抖动要慎用HE 染色的颜色本身携带语义过度抖动会让模型学到错误的颜色先验。归一化用按通道的均值和标准差或者直接除以 255 再减 0.5。如果数据来自不同扫描仪建议做一次直方图匹配或 CLAHE 预处理把亮度拉齐否则模型会把设备差异当成类别差异。import cv2 import numpy as np import torch from torch.utils.data import Dataset class CellDataset(Dataset): def __init__(self, img_dir, mask_dir, size512, trainTrue): self.img_dir, self.mask_dir img_dir, mask_dir self.size, self.train size, train self.names sorted([f for f in os.listdir(img_dir) if f.endswith(.png)]) 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) # OpenCV 默认 BGR转 RGB mask cv2.imread(os.path.join(self.mask_dir, name), cv2.IMREAD_GRAYSCALE) img cv2.resize(img, (self.size, self.size), interpolationcv2.INTER_LINEAR) mask cv2.resize(mask, (self.size, self.size), interpolationcv2.INTER_NEAREST) if self.train and np.random.rand() 0.5: img, mask img[:, ::-1], mask[:, ::-1] # 水平翻转图像和掩膜同步 img img.astype(np.float32) / 255.0 img (img - 0.5) / 0.5 mask (mask 127).astype(np.float32) # 二值化255 - 1 img torch.from_numpy(img).permute(2, 0, 1) mask torch.from_numpy(mask).unsqueeze(0) return img, mask def __len__(self): return len(self.names)逻辑说明读图时 OpenCV 是 BGR必须转 RGB否则和预训练权重或可视化对不上。图像用双线性插值缩放掩膜必须用最近邻否则边缘会出现 0 到 255 之间的中间值二值化后边界错位。翻转时图像和掩膜要同步这是新手最常写错的地方。参数说明size 建议 512显存不够降到 256归一化用 (x-0.5)/0.5 把范围压到 [-1,1]配合 BN 收敛更稳mask 阈值取 127如果标注是 0/1 而非 0/255直接判断大于 0 即可。3.3 损失函数与评价指标Dice 和 BCE 怎么配细胞分割普遍存在类别不平衡背景像素远多于细胞像素单用 BCE 会让模型倾向全预测背景。常见做法是 BCE 和 Dice 按权重相加BCE 稳定梯度Dice 直接优化重叠度。import torch.nn.functional as F def dice_loss(pred, target, eps1e-6): pred torch.sigmoid(pred) inter (pred * target).sum(dim(2, 3)) union pred.sum(dim(2, 3)) target.sum(dim(2, 3)) return 1 - ((2 * inter eps) / (union eps)).mean() def bce_dice(pred, target, bce_w0.5): return bce_w * F.binary_cross_entropy_with_logits(pred, target) \ (1 - bce_w) * dice_loss(pred, target)逻辑说明dice_loss 里先对 logits 做 sigmoid 再算重叠eps 防止分母为零。bce_dice 用权重把两者线性组合bce_w 默认 0.5边界要求高时可以调到 0.3让 Dice 占更大比重。参数说明eps 取 1e-6 足够bce_w 在 0.3 到 0.7 之间调低于 0.3 训练早期容易不稳高于 0.7 又回到类别不平衡问题。评价指标用 Dice 和 IoUDice 对边界更敏感IoU 更严格两个都报才看得出模型真实水平。4. 训练与推理把模型真正跑起来4.1 训练循环与学习率调度训练循环本身不复杂关键是混合精度和梯度裁剪前者省显存后者防梯度爆炸。细胞分割数据量通常不大几十个 epoch 就能收敛。import torch from torch.utils.data import DataLoader from torch.cuda.amp import autocast, GradScaler def train_one_epoch(model, loader, optimizer, scaler, device): model.train() total 0.0 for img, mask in loader: img, mask img.to(device), mask.to(device) optimizer.zero_grad() with autocast(): # 混合精度省显存提速 pred model(img) loss bce_dice(pred, mask) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) # 梯度裁剪 scaler.step(optimizer) scaler.update() total loss.item() * img.size(0) return total / len(loader.dataset) # 优化器与调度器 model UNet(in_ch3, out_ch1, base32).to(device) optimizer torch.optim.AdamW(model.parameters(), lr1e-3, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max50) scaler GradScaler()逻辑说明autocast 把前向和损失计算放到半精度GradScaler 负责缩放梯度避免下溢。clip_grad_norm_ 把梯度范数限制在 1.0细胞分割里 BN 加小 batch 容易梯度尖峰裁剪能明显减少训练崩掉的情况。优化器用 AdamW权重衰减 1e-4学习率 1e-3 起步余弦退火到接近零。参数说明batch size 在 512 输入下 8GB 显存大概能放 4 到 8太小会让 BN 统计不稳可以换 GroupNormT_max 设成总 epoch 数如果 loss 在前几个 epoch 就变成 nan先检查 mask 是否归一化、学习率是否过大。4.2 推理与后处理从概率图到细胞掩膜推理阶段输出的是 logits要经过 sigmoid、阈值化再做连通域过滤去掉小噪点。torch.no_grad() def predict(model, img_tensor, threshold0.5, min_area30): model.eval() prob torch.sigmoid(model(img_tensor.unsqueeze(0).to(device)))[0, 0].cpu().numpy() binary (prob threshold).astype(np.uint8) num, labels, stats, _ cv2.connectedComponentsWithStats(binary, connectivity8) clean np.zeros_like(binary) for i in range(1, num): # 0 是背景从 1 开始 if stats[i, cv2.CC_STAT_AREA] min_area: clean[labels i] 1 return prob, clean逻辑说明sigmoid 把 logits 转成 0 到 1 的概率阈值 0.5 是默认起点。connectedComponentsWithStats 找出所有连通域面积小于 min_area 的当作噪点丢掉这一步对细胞分割很关键模型常在背景里冒出零星假阳性。参数说明threshold 在 0.4 到 0.6 之间调偏高召回低、偏低误检多按任务取舍min_area 按细胞实际像素面积设512 输入下细胞核通常几十到几百像素取 30 是保守值太小去不掉噪点太大会误删小细胞。4.3 训练自己的数据集要改哪几个地方拿到自己的显微镜数据改动集中在三处一是 CellDataset 里的读图方式灰度图把 IMREAD_COLOR 换成 IMREAD_GRAYSCALE模型 in_ch 改成 1二是 mask 的标注值如果标注是 0/1 而不是 0/255二值化阈值相应调整三是输入尺寸细胞特别小的数据集可以把 size 提到 1024但显存要跟上或者用滑窗推理。如果类别不止两类比如要区分细胞核和细胞质把 out_ch 改成类别数损失换成 CrossEntropy 加 Dice 的多类版本mask 读进来保持整数标签不要二值化。这一步改错的表现是 loss 一直不降或者预测全是同一类。5. 避坑与排查细胞分割里最容易翻车的五件事5.1 掩膜和图像没对齐Dice 死活上不去现象训练 loss 能降但验证 Dice 卡在 0.3 左右可视化发现预测掩膜整体偏移几个像素。 原因图像缩放用双线性、掩膜也用双线性或者翻转时只翻了图像没翻掩膜导致两者空间对应关系被破坏。 解决掩膜一律用最近邻插值所有几何变换图像和掩膜同步执行写个可视化脚本把原图和掩膜叠加看一眼对齐问题一眼就能发现。5.2 验证集指标虚高上线就崩现象验证 Dice 0.9换一批新切片掉到 0.5。 原因按图像随机划分同一张切片的相邻图块同时进了训练和验证信息泄漏。 解决按患者或切片 ID 划分同一来源的图块只出现在一个集合里。数据量小时宁可用交叉验证也别图省事随机划分。5.3 训练几个 epoch 后 loss 变 nan现象前几个 epoch 正常突然 loss 变成 nan之后再也降不下来。 原因混合精度下梯度下溢或者学习率过大导致梯度爆炸BN 在小 batch 下统计量不稳也会加剧。 解决加梯度裁剪把 GradScaler 用上学习率从 1e-3 降到 3e-4 试batch 太小就把 BN 换成 GroupNorm细胞分割里这个替换很常见。5.4 预测结果全是背景或全是细胞现象输出掩膜要么全黑要么全白Dice 接近 0 或接近 1 但明显不对。 原因类别不平衡太严重BCE 主导了损失或者 mask 归一化写错标签值不是 0 和 1。 解决确认 mask 二值化正确把 bce_w 调低让 Dice 占主导或者用带 pos_weight 的 BCE。全白通常是阈值太低先把 threshold 提到 0.5 以上看概率图分布。5.5 显存不够512 输入直接 OOM现象一开训练就报 CUDA out of memory。 原因base_channels 64 起步、batch size 开太大、输入尺寸 512 三者叠加。 解决base 降到 32 或 16batch 降到 2 到 4开启混合精度还不行就用梯度累积模拟大 batch。推理阶段用滑窗或把图缩到 256 再上采样回原尺寸精度损失可控。6. 进阶用嵌套深度和深监督把 UNet 调出该有的收益UNet 的嵌套深度不是越深越好。工程里我用一个参数控制中间节点的层数常见取 2 到 3。层数增加会带来两个变化一是参数量和显存上升二是深监督可以接进来。深监督的做法是在每个中间节点输出后接一个 1×1 卷积和上采样单独算一份损失最后加权求和。这样浅层节点也能收到梯度收敛更快边界更准。# 深监督损失示意每个中间输出单独算 Dice再按权重相加 def deep_supervision_loss(outputs, target, weights(1.0, 0.5, 0.25)): loss 0.0 for out, w in zip(outputs, weights): loss w * bce_dice(out, target) return loss逻辑说明outputs 是模型返回的多尺度输出列表从深到浅排列权重递减深层输出主导。权重不是固定的数据量大时浅层权重可以再降数据量小时适当提高帮助收敛。参数说明weights 长度要和 outputs 一致常见三层取 (1.0, 0.5, 0.25)如果验证指标不升反降先把深监督关掉确认基础 UNet 能正常收敛再打开。验证改进是否有效我习惯固定三件事同一份数据划分、同一个随机种子、同一套评价脚本。只改模型结构或损失其他不动跑三次取均值。细胞分割的指标抖动不小单次结果差一两个点说明不了问题三次均值才可信。另外准备一个可视化脚本把原图、真值、预测叠成三栏图指标之外肉眼看一眼边界很多问题指标看不出来图上一目了然。我自己踩得最深的一次是花了两天调 UNet 的嵌套结构最后发现验证集划分泄漏指标全是假的。从那以后我养成的习惯是任何模型改动之前先把数据管线和评价脚本单独验证一遍用一个小数据集过拟合到 Dice 0.99确认整条链路没问题再去谈结构改进。这个习惯帮我省下的时间比任何调参技巧都多。希望帮到你。本文还有配套的精品资源点击获取
返回列表