ARTICLE DETAIL

资讯详情

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

基于Unet的心脏分割实战:Python源码与模型详解

基于Unet的心脏分割实战:Python源码与模型详解 简介本资源面向计算机、人工智能、通信工程等专业的在校学生与教师以及从事医学图像处理的开发者提供一套基于U-Net网络实现心脏分割任务的完整Python源码与训练模型可用于毕业设计、课程设计、作业提交或项目初期立项演示。压缩包共620个文件约53.4MB其中597个png为训练与测试图像样本12个py脚本负责数据加载、模型搭建与训练推理另有2个h5权重文件、若干txt说明及md文档并附带miou-pa-cpa评估结果文件便于直接复现分割效果。资源内代码均经过实际运行测试功能正常已有460人学习下载。读者可据此掌握医学图像语义分割的完整流程包括数据预处理、U-Net结构实现、模型训练与指标评估也可在现有代码基础上修改以适配其他器官分割任务适合作为深度学习入门进阶与实战参考。1. 心脏分割任务为什么值得用 Unet 从头跑一遍心脏磁共振影像的分割是医学影像里少有的「任务定义清晰、数据获取门槛不算高、但效果差距极大」的方向。左心室、右心室、心肌这三类结构在短轴切面上边界模糊、灰度不均还伴随呼吸运动伪影传统阈值法基本没法稳定工作。Unet 这套编码器-解码器加跳跃连接的结构恰好能在小样本医学数据上把像素级分类做到可用水平所以它成了心脏分割任务里最常见的基线模型。这份「基于 Unet 实现的心脏分割任务 python 源码 模型」适合两类人一类是想跑通一个完整医学分割 pipeline 的 python 入门者另一类是手里有自己数据、想拿现成 Unet 结构做迁移的从业者。下面我按「数据怎么准备、模型怎么搭、训练怎么调、坑在哪」的顺序把能直接复现的路径讲清楚。2. 心脏分割数据准备与 Unet 输入输出对齐2.1 心脏短轴数据集的常见组织方式心脏分割公开数据里短轴 cine MRI 一般以「一个病人一个目录、若干切片、每张切片配一张标注 mask」的形式给出。标注通常是 0 背景、1 左心室血池、2 心肌、3 右心室这种整数编码。拿到数据后第一件事不是急着写模型而是把目录结构统一成训练脚本能吃的格式。我一般会整理成下面这样dataset/ train/ patient001_slice01.png patient001_slice01_mask.png patient001_slice02.png patient001_slice02_mask.png val/ ...图像和 mask 必须严格同名对应否则后面 dataloader 会静默错配训练 loss 看着在降实际学的是错位标签这是血泪经验。切片数量少的病人不要直接丢心脏分割本来就样本稀缺可以留作验证集。2.2 灰度归一化与尺寸统一MRI 的灰度范围跟 CT 不一样没有固定的 HU 值所以不能照搬 CT 的窗宽窗位。常见做法是逐切片做 z-score 归一化把每个病人自己的灰度分布拉齐import numpy as np import cv2 def normalize_slice(img): img img.astype(np.float32) mean img.mean() std img.std() 1e-8 img (img - mean) / std # 截断到 [-3, 3]避免个别亮斑把梯度带偏 img np.clip(img, -3, 3) return img def load_pair(img_path, mask_path, size(256, 256)): img cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) mask cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) img cv2.resize(img, size, interpolationcv2.INTER_LINEAR) mask cv2.resize(mask, size, interpolationcv2.INTER_NEAREST) img normalize_slice(img) return img, mask这里有两个参数要盯住size建议统一到 256×256心脏结构在这个分辨率下边界信息够用再大显存吃紧mask 的 resize 必须用INTER_NEAREST用线性插值会插出 1.5 这种不存在的类别后面 one-hot 编码直接崩。归一化后的截断范围[-3, 3]是我试过比较稳的截太狠会丢心肌和血池的对比度。2.3 类别不平衡与 mask 编码心脏分割里背景像素通常占 80% 以上直接算交叉熵会让模型倾向于全预测背景。解决办法有两个方向一是损失函数加权二是把 mask 转成 one-hot 后配合 Dice loss。我一般先把 mask 转 one-hotdef mask_to_onehot(mask, num_classes4): # mask: H x W取值 0..num_classes-1 onehot np.zeros((num_classes, mask.shape[0], mask.shape[1]), dtypenp.float32) for c in range(num_classes): onehot[c] (mask c).astype(np.float32) return onehotnum_classes要和你标注文件的类别数严格一致多一类少一类都会让 loss 计算时维度对不上。转完 one-hot 后训练时用 Dice 交叉熵组合Dice 负责拉回小结构交叉熵稳住整体收敛这个组合在心脏分割上比单用任何一个都稳。3. Unet 网络结构搭建与跳跃连接的关键细节3.1 编码器-解码器通道数怎么定Unet 的经典通道配置是 64-128-256-512-1024但心脏分割数据量通常只有几百到几千张切片直接上 1024 通道容易过拟合。我一般把通道砍到 32-64-128-256最深一层 256 就够。下面是一个可直接用的双卷积块import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.block nn.Sequential( nn.Conv2d(in_ch, out_ch, kernel_size3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, kernel_size3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.block(x)padding1保证卷积后空间尺寸不变这样跳跃连接时编码器和解码器的特征图能直接 concat。biasFalse是因为后面跟了 BatchNorm偏置会被 BN 的均值减掉留着只是浪费参数。BatchNorm 在 batch size 小于 4 的时候统计量不稳如果你显存只够跑 batch 2建议换成 GroupNorm。3.2 下采样、上采样与跳跃连接下采样用最大池化上采样用转置卷积或双线性插值。转置卷积容易产生棋盘格伪影心脏分割这种边界敏感的任务我更倾向双线性插值加一个 1×1 卷积class Up(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.up nn.Upsample(scale_factor2, modebilinear, align_cornersTrue) self.conv DoubleConv(in_ch, out_ch) def forward(self, x1, x2): x1 self.up(x1) # 跳跃连接编码器特征与上采样特征在通道维拼接 x torch.cat([x2, x1], dim1) return self.conv(x)align_cornersTrue在 PyTorch 里对双线性插值的对齐行为有影响设成 True 时角点像素对齐分割边界不会整体偏移半个像素。torch.cat的维度是 1也就是通道维拼接后通道数翻倍所以DoubleConv的in_ch要按拼接后的数量传。跳跃连接是 Unet 能恢复边界细节的核心去掉它分割结果会明显糊掉这一点在心脏心肌这种细结构上尤其明显。3.3 输出层与损失函数对接最后一层用 1×1 卷积把通道数压到类别数不加 softmax因为损失函数里用CrossEntropyLoss或带 logits 的 Dice 更数值稳定class UNet(nn.Module): def __init__(self, in_ch1, num_classes4): super().__init__() self.inc DoubleConv(in_ch, 32) self.down1 nn.MaxPool2d(2) self.conv1 DoubleConv(32, 64) self.down2 nn.MaxPool2d(2) self.conv2 DoubleConv(64, 128) self.down3 nn.MaxPool2d(2) self.conv3 DoubleConv(128, 256) self.up1 Up(256 128, 128) self.up2 Up(128 64, 64) self.up3 Up(64 32, 32) self.out nn.Conv2d(32, num_classes, kernel_size1) def forward(self, x): x1 self.inc(x) x2 self.conv1(self.down1(x1)) x3 self.conv2(self.down2(x2)) x4 self.conv3(self.down3(x3)) x self.up1(x4, x3) x self.up2(x, x2) x self.up3(x, x1) return self.out(x)Up的in_ch写成「上采样特征通道 跳跃特征通道」比如256 128这是最容易写错的地方写错会在torch.cat处直接报维度不匹配。输出层num_classes4对应背景、左心室、心肌、右心室如果你的标注只有三类这里要改成 3否则训练时标签越界。4. 训练循环、评估指标与显存控制4.1 训练脚本的最小可跑版本训练循环里要同时监控 loss 和 Dice光看 loss 会被背景像素带偏。下面是一个精简版训练片段import torch from torch.utils.data import DataLoader device torch.device(cuda if torch.cuda.is_available() else cpu) model UNet(in_ch1, num_classes4).to(device) optimizer torch.optim.Adam(model.parameters(), lr1e-3) ce_loss torch.nn.CrossEntropyLoss() def dice_loss(logits, target, num_classes4): probs torch.softmax(logits, dim1) target_onehot torch.nn.functional.one_hot(target, num_classes).permute(0, 3, 1, 2).float() dims (0, 2, 3) inter torch.sum(probs * target_onehot, dims) union torch.sum(probs target_onehot, dims) dice (2 * inter 1e-6) / (union 1e-6) return 1 - dice.mean() for epoch in range(50): model.train() for img, mask in train_loader: img, mask img.unsqueeze(1).to(device), mask.long().to(device) optimizer.zero_grad() logits model(img) loss ce_loss(logits, mask) dice_loss(logits, mask) loss.backward() optimizer.step()img.unsqueeze(1)是把H×W变成1×H×W因为 Unet 输入要求有通道维。mask.long()是因为CrossEntropyLoss要求标签是 int64。学习率1e-3配 Adam 是常见起点如果 loss 震荡就降到3e-4。Dice loss 里的1e-6是防止分母为零别省。4.2 评估指标与验证集用法心脏分割最常看的指标是每个类别的 Dice 和平均 Dice。验证时不要用训练时的增强数据否则指标虚高。我一般每个 epoch 跑一次验证保存验证集平均 Dice 最高的权重而不是最后一个 epoch 的权重因为过拟合往往在后期def evaluate(model, loader, device, num_classes4): model.eval() dice_sum 0.0 count 0 with torch.no_grad(): for img, mask in loader: img img.unsqueeze(1).to(device) mask mask.long().to(device) logits model(img) pred torch.argmax(logits, dim1) for c in range(1, num_classes): # 跳过背景 p (pred c).float() t (mask c).float() inter (p * t).sum().item() union p.sum().item() t.sum().item() dice_sum (2 * inter 1e-6) / (union 1e-6) count 1 return dice_sum / max(count, 1)跳过背景类别是因为背景 Dice 通常接近 1算进去会把真实结构的分割质量掩盖掉。验证集病人要和训练集病人完全不重叠同一病人的不同切片不能一边训练一边验证否则指标会虚高一大截。4.3 显存不够时的三个调整顺序显存爆了不要一上来就换小模型按这个顺序调先把 batch size 降到 2 或 1再把输入尺寸从 256 降到 192最后才考虑砍通道数。BatchNorm 在 batch 1 时基本失效这时候把nn.BatchNorm2d换成nn.GroupNorm(num_groups8, num_channelsout_ch)训练稳定性会好很多。混合精度训练torch.cuda.amp能省将近一半显存但 Dice loss 里有 softmax建议在 loss 计算时切回 float32避免数值下溢。5. 心脏分割训练里最容易翻车的几个地方5.1 现象loss 一直降但 Dice 不动原因通常是标签错位或类别编码不一致。图像和 mask 文件名没对上、mask resize 用了线性插值、one-hot 类别数和输出层不一致都会造成这个现象。解决方法是训练前先可视化一批img和mask叠加图肉眼确认标签贴在正确结构上再检查num_classes在数据、模型、损失三处是否一致。5.2 现象验证 Dice 高但预测图全是背景这是典型的类别不平衡没处理。交叉熵被背景主导模型学会全预测背景就能拿到很低的 loss。解决方法是把 Dice loss 权重调高或者在交叉熵里给前景类别加权重比如背景权重 0.1、前景权重 1.0。验证时一定要单独看前景类别的 Dice不要只看平均。5.3 现象训练到一半 loss 突然变 NaN多半是学习率太大或者归一化没做好。MRI 个别切片有亮斑z-score 后没截断梯度爆炸。解决方法是加梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)同时确认归一化里有np.clip。如果已经 NaN从上一个保存的权重重启别硬跑。5.4 现象不同病人之间 Dice 波动极大心脏分割对切片位置敏感基底部和心尖的切面结构差异大。如果训练集里某些位置的切片太少模型在那些位置就崩。解决方法是在数据划分时按切片位置分层采样保证训练集覆盖从基底到心尖的完整范围而不是随机按病人切分。5.5 现象推理时单张图比训练时慢很多常见原因是推理没加torch.no_grad()或者模型还在 train 模式导致 BatchNorm 用 batch 统计量。解决方法是推理前model.eval()并包在with torch.no_grad():里这两步能让推理速度和显存占用都明显改善。6. 把 Unet 心脏分割推到可用的两个进阶技巧第一个技巧是测试时增强TTA。心脏分割边界对翻转敏感推理时把原图、水平翻转、垂直翻转各跑一遍把 softmax 概率平均后再取 argmaxDice 通常能涨 1 到 2 个点。代价是推理时间翻三倍适合对精度要求高、对速度不敏感的离线场景。实现上就是把evaluate里的logits换成三次推理的平均def predict_tta(model, img, device): model.eval() with torch.no_grad(): x img.unsqueeze(0).unsqueeze(0).to(device) p1 torch.softmax(model(x), dim1) p2 torch.softmax(torch.flip(model(torch.flip(x, dims[3])), dims[3]), dim1) p3 torch.softmax(torch.flip(model(torch.flip(x, dims[2])), dims[2]), dim1) prob (p1 p2 p3) / 3 return torch.argmax(prob, dim1).squeeze(0)翻转后要把输出翻回来再平均顺序错了概率图对不上结果反而更差。第二个技巧是后处理去小连通域。心脏分割偶尔会在远离心脏的位置冒出几个孤立像素用连通域分析把面积小于阈值的区域归为背景能去掉大部分这类噪声。阈值按你输入尺寸定256×256 下我一般设 30 像素。这两个技巧都不改模型结构属于低成本提点手段。我自己跑心脏分割的习惯是先把基线 Dice 跑出来记在本子上再逐个加 TTA 和后处理每次只改一个变量确认涨点再留。这样即使某个技巧在你数据上不work也能快速定位是数据问题还是技巧问题。希望帮到你。本文还有配套的精品资源点击获取
返回列表