尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

基于CNN联合学习的心室分割与心脏病分类:可复现源码解析

基于CNN联合学习的心室分割与心脏病分类:可复现源码解析 简介面向医学影像分析与深度学习实践者压缩包提供了一套基于CNN的联合学习方案将心室分割与心脏病分类任务统一建模适合需要复现多任务学习实验或开展相关课题的研究生与工程师。包内共29个文件以Python源码为绝对主体涵盖网络结构搭建、Dice与MS-SSIM损失设计、数据加载、独立分割训练/测试、分割分类联合训练与推理检测等完整流程附带PDF实验报告详细介绍早期网络设计对比另有README说明与预训练权重文件。压缩包整体约30.6MB结构按模块分层便于直接对照目录学习。目前已有93人学习下载。读者可获得一套可直接运行的多任务分割分类项目包含独立与联合两种训练范式、损失函数自定义实现、测试检测工具以及实验报告中对网络设计取舍的分析能有效节省从零搭建环境与调参的时间。1. 基于CNN联合学习的心室分割与心脏病分类一份能被完整复现的项目源码心室边界和心脏病的类型看起来是两件事一个要逐像素把心室轮廓抠出来一个要整张图判断病理归属。这套基于CNN联合学习的项目源码把这两件事放进同一个网络——共享一套编码器提取特征分割分支输出心室区域的概率图分类分支输出疾病类别以ACDC数据集为例就是NOR、MINF、DCM、HCM四类训练完成后还有一份实验报告把Loss曲线、Dice、Accuracy和结论都写清楚。对刚接触医学影像深度学习的工程师来说它是少有的、从数据到训练再到评估全部闭环的参考实现对想迁移到其他数据集做实验的人来说改分类分支和损失权重就能跑自己的场景。下面按原理、数据流、训练参数、翻车点和验证手段依次拆开讲。2. 网络结构拆解共享编码器、双头输出与联合损失设计2.1 联合学习解决什么问题三个实际理由分开做分割和分类会得到两个模型、两套特征。但心脏MRI上判断心肌病的依据恰恰是心室壁厚度、心腔面积、局部运动异常这些空间特征——分割模型在逐像素标注过程中被迫学会了这些特征分类模型却要从头再学一遍。共享编码器之后两个任务共用同一份底层特征分类分支相当于白拿了分割分支的边界表征这就是联合学习在医学影像小数据下比两个单任务模型更稳的根本原因。第二个理由是标签利用更充分。分类标签每个病例都有而分割标签需要逐层手工勾画成本高、数量稀缺。联合训练时分类分支的梯度也能反向传播到共享编码器相当于用标注充分的分类任务辅助标注稀疏的分割任务。在小样本场景100例级别下这本身就像一种隐式正则能明显压低分割模型在验证集上的过拟合。第三个理由很朴素但工程上很实用一次forward同时拿到分割mask和类别概率推理端不需要维护两套模型、两套预处理管线。这也是这类项目普遍存在的价值点。2.2 一个通用的JointUNet结构编码器4层下采样双头各取所需项目里常见做法是直接沿用UNet的编码器作为共享主干后面接两个头。输入是单通道MRI切片尺寸一般定在256×256。编码器走4个stage每个stage是两层Conv3x3BNReLU再接MaxPool通道数按32→64→128→256递增最后得到空间分辨率缩小4倍的特征图。分割头用UNet右侧的解码器结构三次上采样每一层把编码器同尺度的输出concat进来恢复空间细节最后用1x1卷积输出2通道背景、心室接Softmax。分类头则不一样它不再上采样而是直接在最后一层特征图上做全局平均池化然后接FC256→FC4。用GAP而不是直接Flatten是为了保留一点点空间位置信息的同时抑制全连接层的过拟合在100例规模的数据集上非常实用。表JointUNet各模块的输入输出模块输入输出具体构成共享编码器1×256×256256通道、64×64特征图4个stage每stage两层Conv3x3BNReLUMaxPool通道32→64→128→256分割头256特征图 各层skip连接2×256×256概率图三次上采样concat最后1x1卷积出2通道分类头256通道特征图4类logitsGAP→FC256→FC4分割头需要浅层特征来精修边界分类头则更依赖深层语义。这种结构安排下浅层位置的卷积共享给分割深层位置的卷积偏向分类各自的梯度在反向传播时天然完成了分工这也是它比简单“UNet后面直接接FC”更稳的原因。2.3 联合损失函数Dice加CE权重按量级调分割任务不能只用CrossEntropy因为心室区域在整张图里占比通常不到10%类别极端不平衡会让模型一股脑全预测成背景。常见做法是Dice Loss和CE按比例混合或者直接单用DiceLoss。分类任务用标准CrossEntropyLoss。联合损失就是两个loss按权重相加L λ1 × L_seg λ2 × L_cls。这里最需要盯的是两个loss的量级差异。DiceLoss的值通常在0.1到0.9之间CE则可能小到1e-3直接相加时分割任务会把梯度和训练方向完全吞掉。我一般先用均衡配置跑10个epoch看曲线再根据两个loss的实际量级调整。如果分类loss几乎不下降就把λ2往上抬。表不同任务倾向下的权重参考目标场景λ1分割λ2分类适用情况分割为主1.00.1~0.3只关心心室边界精度两者均衡0.80.5项目默认的稳妥起点分类为主0.51.0临床筛查场景更看重类别判断这个权重不是玄学它会直接决定两个分支的收敛速度差。后面避坑章节里我会把“分类acc被分割任务压死”的现象和调参流程单独展开。3. 数据流水线MRI预处理、mask对齐与训练样本组织3.1 ACDC这类公开数据集的组织方式与常见坑项目的数据来源大概率是ACDCAutomated Cardiac Diagnosis Challenge这类公开心脏MRI数据集。它的组织形式是每个病人一个文件夹里面有短轴位的DICOM序列3D的心脏MRI切片堆叠以及对应的.nii.gz分割标签文件。分割标签通常是三类背景、右心室、左心室腔和心肌具体取决于源码里的类别定义。新手最容易栽在第一步DICOM和NIfTI mask的物理坐标不一致。DICOM序列读取后有一个Image Orientation和OriginNIfTI mask也有自己的Direction二者如果没对齐就重采样出来的图像和mask是错位的。所以预处理流程里必须先统一坐标系再做任何裁剪或缩放。这步做错了后面训练Dice再高都是假的。3.2 预处理流水线重采样、z-score归一化、中心裁剪我一般会写一个和下面类似的预处理脚本把每个病人处理成npy文件训练时直接加载省得每次都在线读DICOM。import os import numpy as np import SimpleITK as sitk def preprocess_patient(dicom_dir, mask_path, out_dir, target_size256, spacing(1.0, 1.0, 1.0)): # 1) 读取DICOM序列得到3D MRI体数据 reader sitk.ImageSeriesReader() dicom_names reader.GetGDCMSeriesFileNames(dicom_dir) reader.SetFileNames(dicom_names) image reader.Execute() mask sitk.ReadImage(mask_path) # 2) 重采样到统一物理间距图像用线性插值mask用最近邻 img_resampler sitk.ResampleImageFilter() img_resampler.SetOutputSpacing(spacing) img_resampler.SetInterpolator(sitk.sitkLinear) image_r img_resampler.Execute(image) mask_resampler sitk.ResampleImageFilter() mask_resampler.SetOutputSpacing(spacing) mask_resampler.SetInterpolator(sitk.sitkNearestNeighbor) mask_r mask_resampler.Execute(mask) # 3) 取数组并逐层切片处理 arr sitk.GetArrayFromImage(image_r) # (D, H, W) lab sitk.GetArrayFromImage(mask_r) # (D, H, W) D, H, W arr.shape os.makedirs(out_dir, exist_okTrue) for i in range(D): slice_2d arr[i].astype(np.float32) # z-score归一化比全局min-max更抗高亮噪声 mean_v, std_v slice_2d.mean(), slice_2d.std() 1e-8 slice_2d (slice_2d - mean_v) / std_v # 中心裁剪出心脏区域降低背景占比稳定Dice训练 h0 max((H - target_size) // 2, 0) w0 max((W - target_size) // 2, 0) slice_crop slice_2d[h0:h0target_size, w0:w0target_size] mask_crop lab[i][h0:h0target_size, w0:w0target_size] np.save(os.path.join(out_dir, fimg_{i:03d}.npy), slice_crop) np.save(os.path.join(out_dir, fmask_{i:03d}.npy), mask_crop) print(fprocessed {D} slices for {os.path.basename(dicom_dir)})逻辑说明第2步统一物理间距是为了让不同扫描设备和不同切片厚度的数据在进入网络前保持一致的尺度图像用线性插值保证灰度过渡自然mask用最近邻插值防止类别边界被平滑掉。第3步z-score归一化是针对MRI强度没有绝对物理单位的特性用每张切片的均值和标准差做标准化比在整个体数据上做min-max更稳能避免某一层高亮信号把其他层的对比度压没。参数说明target_size由网络输入决定能不能从512裁到256得看你的显存spacing(1.0,1.0,1.0)是把体素重采样成1毫米各向同性这在心脏MRI里比较常用。如果硬件紧张可以放宽到1.5毫米分割精度会轻微下降但训练显存压力小很多。还有一点要注意重采样之后务必打印image_r.GetDirection()和mask_r.GetDirection()确认方向一致再继续否则后面所有切片都是错的。3.3 数据集划分按病人分不能按切片分增强要谨慎医学影像数据有个铁律同一病人的所有切片不能同时出现在训练集和验证集。心脏MRI一次扫描有十几张slice相邻slice高度相似如果按slice随机划分验证集会混入训练病人的近邻切片Dice虚高得离谱换到新病人直接崩。实际做法是按病人列表划分常见比例7:1:2或者直接做5折交叉验证。项目源码的实验报告里如果给出了稳定的指标方差多半就是按病人折的结果。数据增强方面旋转、缩放、亮度扰动是安全的但水平翻转要慎重。心脏的左右心室位置和解剖形态并不是完全对称的有些病变本身就体现在某一侧心室的变化上翻转等于制造了错误监督。我一般只用±10°旋转、0.9到1.1倍缩放、±15%亮度扰动而且mask用相同的仿射参数同步变换。如果项目里用albumentations实现增强记得把mask也传入同一个transform pipeline别只增强图像。4. 训练复现超参数、loss权重、评估指标照着抄4.1 训练入口与核心超参数基线拿到这套源码后我想你第一件事跟我一样不急着改代码先按默认参数跑通一遍确认显卡能跑起来、Loss在降、指标有个基线再动任何优化。训练脚本的核心部分大致长这样。import torch from torch import nn model JointUNet(in_channels1, n_classes4) model.cuda() optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-5) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemin, factor0.5, patience6 ) # 分割用Dice CE混合背景权重0.2、心室权重0.8 loss_seg DiceWithCE(weighttorch.tensor([0.2, 0.8])) loss_cls nn.CrossEntropyLoss() for epoch in range(1, 101): model.train() for x, y_seg, y_cls in train_loader: x x.cuda() y_seg y_seg.cuda() y_cls y_cls.cuda() seg_out, cls_out model(x) # 联合损失分割带1.0权重分类带0.3权重 loss 1.0 * loss_seg(seg_out, y_seg) 0.3 * loss_cls(cls_out, y_cls) optimizer.zero_grad() loss.backward() optimizer.step() if val_score best_val_score: torch.save(model.state_dict(), best_epoch.pth)逻辑说明AdamW配lr1e-4是医学影像分割训练里比较稳的起点weight_decay1e-5给一点正则100例级别的数据集上很有用。y_seg的标签形状有两种常见约定一种是(B, H, W)整数标签另一种是(B, 2, H, W)的one-hot格式具体看DiceWithCE的实现。如果项目里给的mask是整数标签而你的loss期望one-hot记得在dataset的__getitem__里转换别在loss里临时弄。参数说明DiceWithCE(weight[0.2, 0.8])的weight是给CE分支用的类别权重心室只占图像一小块给它0.8的权重能让模型不偷懒。lr1e-4如果遇到Loss不降可以试着调成5e-5或3e-4但不要一上来就动另一个参数一次只变一个否则翻车了根本不知道是谁的锅。4.2 训练过程怎么看先看曲线再下结论训练开始后我习惯同时开TensorBoard和终端日志。tensorboard --logdir runs启动后重点盯三组曲线联合loss、分割Dice、分类Accuracy。正常的情况下前20个epoch里Dice会从0.1往上爬到0.6左右联合loss稳步下降分类Accuracy曲线会有抖动但趋势应该向上。如果看到联合loss下降但val Dice停滞在0.5上下多半是DiceLoss和CE混合出了问题。如果val Dice上去了、但分类Accuracy始终在50%以下多半是两个loss的权重配比不对分类分支的梯度被分割分支盖住了。这时候回到第2章的权重表把λ1:λ2从1.0:0.3改成0.8:0.5重新跑10个epoch看曲线变化。不要等100个epoch跑完再回头分析那不叫训练叫碰运气。4.3 评估指标与项目实验报告对应关系项目源码的评估脚本一般会计算两套指标分割用Dice和IoU分类用Accuracy、F1和混淆矩阵。Dice的定义是2×TP除以2×TP加FP加FNIoU是TP除以TP加FP加FN两者都按slice聚合最后对验证集所有病人取平均。分类部分用sklearn的confusion_matrix直接出四分类矩阵从矩阵里能看出发散型心肌病DCM和正常NOR之间是不是有系统性误判。表分割模型常用指标及其参考区间指标计算方式参考区间Dice2A∩BIoUA∩BAccuracy正确分类样本/全部样本0.85以上算合格F1 Score2×precision×recall/(precisionrecall)0.8以上算合格这些参考值是医学影像分割任务里的常见共识不是这套源码的实测值。拿到实验报告时应该重点看它有没有分门别类地报告每类病人的Dice而不是只给一个总体均值——心肌梗死患者的Dice普遍比正常组低这是正常的现象如果报告里写清楚了这个差异说明实验部分做得比较扎实。评估脚本里如果同时导出了预测mask的图建议人工抽查十张尤其看心室边缘是不是齐整、有没有把右心室当成左心室。5. 避坑清单联合学习在心室分割上的4个常见翻车点5.1 分割Loss不降Dice卡在0.5附近转圈现象训练了30个epochDiceLoss始终在0.6上下浮动val Dice不升反降预测mask一片黑。原因类别不平衡加上归一化不当。心室区域占全图比例本来就小如果预处理用的是全局min-max归一化一旦某个slice左下角有高亮噪声整个图像的灰度会被压扁心室和背景的对比度直接消失。模型想学都学不到东西。解决回到第3章的预处理把min-max换成z-score并做中心裁剪把背景区域压下去。同时检查DiceWithCE的weight参数给心室类加权重。改完这两处Dice在前5个epoch就应该从0.1快速爬到0.6以上否则继续检查数据不是模型的问题。5.2 联合训练后分类Accuracy反而不如单独分类现象单独训练分类分支时Accuracy有0.78联合训练后掉到0.62Loss曲线里分类部分一直高频抖动。原因两个loss的量级没对齐。DiceLoss在0.1到0.9之间CE可能只有0.01到0.5默认的λ配置让分割梯度过大分类头的梯度在共享编码器里被抵消相当于分类分支一直在学但底层特征完全被分割任务重塑。解决按第2章的权重表先给分类提权从1.0:0.3改成0.8:0.5重新跑。如果还不行就做分离学习率编码器和分割头走lr1e-4分类头单独走lr3e-4让分类头在上层收敛得更快一些。这个操作在PyTorch里把分类头参数单独放进一个param group就行。5.3 验证集Dice高但目视预测mask严重错位现象Dice报告0.89但把预测mask叠加到MRI上发现心室轮廓整体向一侧偏移了几个像素或者把右心室也圈了进去。原因mask和图像的重采样方向不一致导致标签和图像之间有几个像素的平移误差。Dice对微小的整体偏移并不敏感它奖励的是重叠面积所以数值还能维持在高位。目视检查才是拆穿这个问题的唯一办法。解决在预处理脚本里重采样后用Print打印图像和mask的Origin、Direction、Spacing三个属性确认物理坐标完全一致。如果发现方向不一致需要用sitk.Resample统一到同一个参考坐标系而不是直接转numpy。从那以后我每次处理完一个病人的数据都强制用matplotlib叠画五张图检查一下再进入训练队列。5.4 换一台机器训练直接OOM现象源码备注的是8G显存可训换到4G的卡上batch_size8的默认参数跑不起来报CUDA out of memory。原因输入图像是256×256没错但模型训练时中间特征图和one-hot标签都很吃显存。特别是seg头每层都concat了编码器同尺度特征显存占用比纯分类模型高一大截。解决batch_size从8减到4输入尺寸从256降到224torch.cuda.amp混合精度打开三个动作一起做显存占用能压到原来的四分之一。如果配置里有gradient checkpointing在编码器的每个stage包一层torch.utils.checkpoint打开也能省30%显存代价是训练速度慢一点。6. 进阶验证用Grad-CAM解释分类依据并导出可部署的单图推理脚本6.1 用Grad-CAM确认分类分支到底在看哪里训练收敛之后模型在验证集上的指标达标只代表结果正确不代表过程正确。医学影像模型最常见的翻车方式是“看对了数据看错了位置”。心脏MRI的边缘经常有扫描床伪影、体外高亮区域分类分支如果只学了这些背景信号Accuracy也能很高但换一台设备的数据直接失效。所以联合学习结束后我必做的一步是Grad-CAM热力图检查。import torch import numpy as np def gradcam(model, x, target_layer): feature_map None gradient None def forward_hook(module, input, output): nonlocal feature_map feature_map output def backward_hook(module, grad_input, grad_output): nonlocal gradient gradient grad_output[0] handle_f target_layer.register_forward_hook(forward_hook) handle_b target_layer.register_full_backward_hook(backward_hook) model.eval() x x.cuda() _, cls_out model(x) target torch.argmax(cls_out, dim1).item() model.zero_grad() cls_out[0, target].backward() handle_f.remove() handle_b.remove() # 对梯度做空间平均再和特征图加权求和 weights gradient.mean(dim(2, 3), keepdimTrue) cam (weights * feature_map).sum(dim1).relu() cam cam.squeeze().cpu().numpy() cam (cam - cam.min()) / (cam.max() - cam.min() 1e-8) return cam逻辑说明Grad-CAM的核心思想是看分类分支对目标类别的梯度集中在特征图的哪些位置。backward_hook拿到梯度后对空间维度做平均得到每个通道对目标类别的贡献权重再和特征图相乘求和最后ReLU只保留正向贡献区域。热力图亮度高的区域就是模型的决策依据。参数说明target_layer建议选分类分支的GAP之前的最后一个卷积层输出而不是最深的共享编码器特征那样分辨率太低。医学影像上热力图应该落在心室壁和心腔区域如果高亮区域跑到图像角落说明模型在偷懒直接回测数据预处理。6.2 单图推理脚本把训练好的模型变成一个能用的工具训练完、验证完最后一步是把模型固定下来写一个单图推理脚本。这个脚本的意义不只是部署更是日常复查的手段——以后每次拿到一个新病人的DICOM跑一遍就能同时看到分割mask和分类概率。import torch import numpy as np import SimpleITK as sitk def inference(dicom_dir, model, devicecpu): reader sitk.ImageSeriesReader() reader.SetFileNames(reader.GetGDCMSeriesFileNames(dicom_dir)) image reader.Execute() # 和训练保持一致重采样、z-score归一化、中心裁剪到256 image sitk.Resample(image, [256, 256, image.GetDepth()], sitk.Transform(), sitk.sitkLinear, [1.0, 1.0, 1.0], image.GetOrigin(), image.GetDirection()) arr sitk.GetArrayFromImage(image).astype(np.float32) x (arr[-1] - arr[-1].mean()) / (arr[-1].std() 1e-8) x torch.from_numpy(x).unsqueeze(0).unsqueeze(0).float() model.eval() with torch.no_grad(): seg_out, cls_out model(x) seg_mask seg_out.argmax(dim1).squeeze().numpy() probs cls_out.softmax(dim1).squeeze().numpy() np.save(pred_mask.npy, seg_mask) np.save(pred_cls.npy, probs) print(f分类概率: NOR{probs[0]:.3f}, MINF{probs[1]:.3f}, fDCM{probs[2]:.3f}, HCM{probs[3]:.3f}) return seg_mask, probs逻辑说明推理脚本和训练数据处理的每一条规则都必须完全一致——同样的重采样间距、同样的z-score、同样的中心裁剪位置任何一个参数对不上模型输出就是不可信的。代码里取切片的第arr[-1]层只是一个示例实际项目里需要根据源码约定的slice索引来。参数说明模型加载时一定要用map_locationcpu再to(device)否则会在没有GPU的机器上报错。保存的seg_mask是整数标签0和1可以用cv2.imwrite叠到原图上做目视检查。如果项目源码里带GUI或可视化工具把这两个npy接进去就是一个完整的病案辅助分析小工具。我自己最早训这类模型的时候只看Dice和Accuracy两个数字直到一次把Grad-CAM热力图叠回原图才发现模型盯的是扫描床伪影而不是心室壁。从那以后每次联合学习实验我都强制走一遍“Grad-CAM 单图推理 目视叠图”三连确认模型真的在看心脏区域才敢把结果写进实验报告。希望帮到你。本文还有配套的精品资源点击获取
返回列表