尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

医学影像分割实战:2.5D/3D U-Net与V-Net完整流程

发布时间:2026/9/14 23:48:24

资讯中心
01
ARTICLE

医学影像分割实战:2.5D/3D U-Net与V-Net完整流程

医学影像分割实战:2.5D/3D U-Net与V-Net完整流程
简介面向Python基础较好的医学影像学习者这套代码资源聚焦基于深度学习的医学影像图像分割适合完成课程设计或入门UNet系列模型。包内提供从数据准备、预处理、模型搭建到训练、预测、后处理的完整流程包含2.5D、3D等不同输入维度的UNet实现并配有损失函数、数据IO等公共模块便于对照理解不同网络结构。资源共19个Python文件压缩包大小仅35KB轻量但结构清晰。文件说明明确划分了训练集、验证集与测试集并给出运行顺序附带的预测结果可辅助检验模型效果。已有551人学习下载适合作为医学影像分割入门或课程设计的直接参考。1. 为什么医学影像分割要改用2.5D/3D模型而不是逐层切2D医学影像分割不是普通自然图像分割的搬运。MRI的体数据是三维的相邻切片之间包含连续解剖结构信息如果一张张切出来用2D U-Net跑等于把立体信息压扁边缘和体积估计都会失真。而直接用3D卷积网络对显卡显存又极不友好一块12G的卡甚至放不下一个稍大体积的批大小。这个项目里同时给出了unet_25d.py、vnet_3d_nn.py等多套网络就是在“信息完整性”和“显存成本”之间做权衡——先用2.5D的切片堆叠把Z轴信息塞进输入通道再提供完整的3D V-Net作为高配方案同时把数据准备、训练、预测、后处理串成一条可复现的流水线。对正在做医学影像课程设计或入门3D分割的工程师来说这套代码可以直接改改路径就能跑也能一步步看清每个环节的输入输出在哪里。2. 从NIfTI到训练样本数据准备与预处理在训练任何分割模型之前需要先把原始nii文件转成模型能吃到的数组。项目里data/train是训练集其中10%留作验证集data/rest是测试集data/test是课程设计要求预测的数据。你要先跑create_train_data.py或create_train_data_25d.py来生成对应的h5/npy中间格式。2.1 读取NIfTI的姿势医学影像文件通常是.nii或.nii.gz用SimpleITK或nibabel读取。两种库都行我推荐nibabel因为它对单模态MRI足够轻。import nibabel as nib import numpy as np img nib.load(data/train/image/case01.nii) volume img.get_fdata().astype(np.float32) label nib.load(data/train/label/case01.nii).get_fdata().astype(np.uint8) print(volume.shape, label.shape) # 例如 (256, 256, 80)这段代码把图像和标签读成numpy数组。get_fdata()会返回体素值label是逐体素的分割类别0通常是背景1、2等是器官。要注意NIfTI的axis顺序是(x, y, z)很多2D可视化工具显示的是z方向切片。2.2 预处理裁剪和归一化MRI的原始体素值范围很大不同扫描参数下灰度分布不一致直接喂网络会崩。preprocess_25d.py里做的工作通常是这几步去掉背景多余的空白区域把非零体素做z-score归一化再统一resolutions/spacing到如果数据来自不同设备。常见的实现方式如下def preprocess_volume(volume, lower_percent0.0, upper_percent99.5): pixels volume[volume 0] lower np.percentile(pixels, lower_percent) upper np.percentile(pixels, upper_percent) volume np.clip(volume, lower, upper) volume (volume - volume.min()) / (volume.max() - volume.min() 1e-8) return volume这里用百分位裁剪去掉极端高信号比如脂肪、空气强度异常然后做min-max归一化把体素压到[0,1]。对标签则不能用插值只能做最近邻重采样否则类别会被平均出小数。2.3 2.5D输入怎么生成2.5D模型输入的不是单张切片而是在当前切片位置上下各取N张切片一起作为通道。这样既保留了部分层间上下文又让网络可以复用2D U-Net的结构。create_train_data_25d.py里会遍历volume的每个切片位置为每个位置生成一个(num_slices, H, W)的张量。def extract_25d_slices(volume, label, slice_num4): # slice_num 表示上下各取多少张总通道数2*slice_num1 H, W, D volume.shape middle slice_num 1 images, labels [], [] for d in range(D): idx np.clip(np.arange(d - slice_num, d slice_num 1), 0, D - 1) multi_slice volume[:, :, idx] # (H, W, 2*slice_num1) images.append(np.transpose(multi_slice, (2, 0, 1)).astype(np.float32)) labels.append(label[:, :, d].astype(np.int64)) return np.stack(images), np.stack(labels)参数slice_num是超参数通常取2~4。取值太小Z轴上下文不足取值太大离当前层太远的切片是噪声。实验下来腹部MRI里slice_num3效果较好对显存占用比3D卷积小一个量级。2.4 数据生成器与数据增强数据量不够时生成器里可以嵌入在线增强随机旋转、翻转、弹性形变等。generator_25d.py里我习惯写成一个继承keras.utils.Sequence的对象每轮打乱索引并实时增强。下表总结了数据准备阶段各脚本的职责运行时请对照自己的文件命名修改路径脚本输入输出说明preprocess_25d.py原始nii归一化后的npy/image字段负责裁剪、归一化、重采样create_train_data.py / create_train_data_25d.pyniilabel2D/2.5D npz/h5文件把体数据切成2D切片或生成多通道切片generator_25d.pynpz/h5文件批数据含增强供训练循环按batch读取pathvariable.py--集中管理数据路径和常量提示如果训练时发现loss不降先回去检查create_train_data生成的数据里label是否对齐常见问题是在重采样时对标签用了线性插值结果每个体素都变成0.几。3. 网络选型2.5D U-Net、3D U-Net与V-Net的取舍光有数据还不够网络骨架决定模型上限。这个项目里包含了unet_25d.py、v_net_25d.py、unet_3d_nn.py和vnet_3d_nn.py。理解它们的区别才知道该跑哪个。3.1 2.5D U-Net用通道换上下文unet_25d.py是最容易上手的版本。它的本质还是2D U-Net只是把输入通道数从1变成2*slice_num1。编码器第一层用3x3卷积把多通道合并后面的结构完全和2D一致。import torch.nn as nn class UNet2_5D(nn.Module): def __init__(self, in_channels7, n_classes2): super().__init__() self.enc1 nn.Sequential( nn.Conv2d(in_channels, 64, 3, padding1), nn.BatchNorm2d(64), nn.ReLU(inplaceTrue), nn.Conv2d(64, 64, 3, padding1), nn.BatchNorm2d(64), nn.ReLU(inplaceTrue)) self.pool1 nn.MaxPool2d(2) # ... 后续层类似这里省略 self.dec4 nn.ConvTranspose2d(512, 256, 2, stride2) self.out nn.Conv2d(64, n_classes, 1) def forward(self, x): x1 self.enc1(x) x self.pool1(x1) # ... 前向传播 return self.out(x)这里in_channels对应2.5D的通道数如果slice_num3就是7。注意第一层卷积没有把Z轴压缩只是加权融合相邻切片因此2.5D对Z轴的分辨率不敏感但显存比纯3D小得多。3.2 3D U-Net与3D V-Netunet_3d_nn.py把所有卷积换成nn.Conv3d输入是(1, D, H, W)或者patch。vnet_3d_nn.py则是V-Net它的特点是使用了残差连接和基于Dice的损失函数而且编码路径使用下采样和卷积同时进行的block参数量比普通3D U-Net更高效。class VNet3D(nn.Module): def __init__(self, in_channels1, n_classes2): super().__init__() self.conv_init nn.Conv3d(in_channels, 16, 3, padding1) self.res nn.Sequential( nn.Conv3d(16, 16, 3, padding1), nn.ReLU(inplaceTrue), nn.Conv3d(16, 16, 3, padding1)) self.out nn.Conv3d(16, n_classes, 1) def forward(self, x): x self.conv_init(x) x x self.res(x) # 残差连接 return self.out(x)残差连接能避免深层网络的梯度消失在医学小数据集上特别管用。3D模型需要输入patch。通常要把体积裁剪成64,64,64或128,128,64的patch训练再在推断时滑动窗口。3.3 怎么选显存、病变尺寸和数据量我的经验是方案显存占用层间信息适用场景2D U-Net低无切片间无关联或数据量极小2.5D U-Net中局部缺显存又想利用层间信息3D U-Net / V-Net高全局精细结构、病变较小、显存≥16G如果你要分割肝脏、肾脏这类器官2.5D基本够用如果分割肿瘤或细小血管3D模型的连续感受野更重要。可以先跑v_net_25d.py它介于2.5D与3D之间把3D卷积作用在小块patch上。4. 训练与验证损失函数、学习率与运行命令拿到数据和模型之后核心问题是怎么把它训练出来。train_25d.py和train.py里包含的流程是加载生成器 → 定义模型 → 选择损失 → 优化器 → 迭代 → 每若干epoch在验证集上评估 → 保存权重。4.1 损失函数的选择医学分割最常见的损失组合是CrossEntropyLoss DiceLoss。纯CrossEntropy在背景占90%以上时会让网络倾向于全预测背景Dice Loss则直接优化类别重合度。项目里的loss_function.py一般会实现这两种下面给出一个带平滑系数的Dice Loss实现import torch import torch.nn.functional as F class DiceLoss(nn.Module): def __init__(self, smooth1.0): super().__init__() self.smooth smooth def forward(self, pred, target): # pred: (B,C,H,W) target:(B,H,W) pred F.softmax(pred, dim1) target_onehot F.one_hot(target, num_classespred.size(1)).permute(0, 3, 1, 2).float() B, C pred.shape[:2] dice 0.0 for c in range(C): p pred[:, c].contiguous().view(B, -1) t target_onehot[:, c].contiguous().view(B, -1) intersection (p * t).sum(dim1) dice (2 * intersection self.smooth) / (p.sum(dim1) t.sum(dim1) self.smooth) return 1 - dice.mean()smooth用于防止分母为0一般取1.0。如果你的训练只有两个类别背景和器官就把num_classes设为2。注意one_hot要和pred的尺寸保持一致维度顺序错误是初学者最容易犯的错。4.2 训练流程train.py的运行逻辑大致是cd /path/to/unet pip install -r environment.txt python3 train_25d.py --data data/train --val_ratio 0.1 --batch_size 4 --epochs 100 --lr 1e-3建议在命令行里明确写参数不要改死路径。val_ratio0.1正是项目里“其中10%为验证集”的体现。训练过程中每个epoch后计算验证集Dice保存最佳模型。逻辑说明train_25d.py会从pathvariable.py读取数据目录遍历train文件夹里的npz文件按val_ratio切出验证集。优化器我建议先用Adamlr1e-3训练到中段再换成SGDmomentum继续调优这样收敛快且后期更稳。4.3 显存不够时怎么办不要一次性读入整个volume用patch训练patch大小设为64,64,32之类。用混合精度训练PyTorch下用torch.cuda.amp.autocast()。减小slice_num从3改成1显存立刻减半。用梯度累积每4个小batch累计一次梯度更新等效于增大batch_size。4.4 训练时看哪些指标除了loss必须记录验证集的Dice、IoU和表面距离如果数据有边界。loss曲线只能看收敛趋势Dice才反映分割质量。我在实际调试中发现如果训练集Dice涨到0.9而验证集停在0.7多半是过拟合或数据预处理时标签和图像未对齐而不是网络表达能力不够。5. 预测与后处理把分割结果落成文件训练完成后predict.py和predict_rest.py负责对data/test/image和data/rest/image做推断。这步看似简单实际有坑模型输入要经过和训练一致的预处理输出概率图要还原回原体积尺寸后处理决定最终掩码质量。5.1 单文件预测流程以predict.py为例它对data/test/image下每个nii文件预测并保存到data/test/predict。代码核心如下import SimpleITK as sitk import numpy as np import torch def predict_volume(model, volume_path, save_dir): img sitk.ReadImage(volume_path) volume sitk.GetArrayFromImage(img).astype(np.float32) # (z,y,x) volume preprocess_volume(volume) input_tensor torch.from_numpy(volume[None, None]).float().cuda() with torch.no_grad(): logits model(input_tensor) prob torch.softmax(logits, dim1).cpu().numpy()[0, 1] # 取前景类 out (prob 0.5).astype(np.uint8) result_img sitk.GetImageFromArray(out) result_img.CopyInformation(img) sitk.WriteImage(result_img, f{save_dir}/{os.path.basename(volume_path)})注意SimpleITK读进来是(z,y,x)顺序模型训练时如果用的是(x,y,z)这里要转置。CopyInformation会保留原始spacing和origin否则后续医学软件无法正确读取。5.2 三类文件与三次运行项目里data/rest/predict、predict1、predict11是三次运行结果说明predict_rest.py被跑过多次。为什么要跑三次因为如果推理过程没有设置固定随机种子或使用了dropout每次predict结果会有细微波动。对课程设计来说重复跑三次把三个结果取多数投票能稳定最终指标。这也是一个实用小技巧可以写一个bagging式的投票函数def majority_vote(prob_list): # prob_list: [(H,W,D)]*3 stack np.stack(prob_list, axis0) avg stack.mean(axis0) return avg 0.55.3 后处理去小连通域和填洞postprocess.py里主要做两件事删除体积小于阈值的连通域消除噪声预测填补前景内部的小孔让分割更完整。用scipy.ndimage即可from scipy import ndimage as ndi def remove_small_objects(mask, min_size500): mask ndi.binary_opening(mask, iterations1) labeled, num ndi.label(mask) sizes ndi.sum(mask, labeled, range(1, num 1)) keep np.where(sizes min_size)[0] 1 filtered np.isin(labeled, keep) return filteredmin_size需要根据体素spacing换算比如目标是去掉小于5mm³的噪声体素体积为1.5mm³那min_size3。不是所有数据集都适合统一阈值先观察预测结果再调。6. 进阶技巧在数据不足下提升分割稳定性最后我分享一下在不额外采集数据的前提下让现有分割模型更稳的几个具体做法。这些方法在这个课程设计项目里可以直接套用不依赖额外GPU资源。6.1 使用冻结预训练编码器如果你的2.5D U-Net编码器换成在ImageNet上预训练的ResNet34把输入通道改成多通道后第一层卷积权重用平均值初始化前20个epoch冻结编码器只训练解码器。对于MRI这种灰度图像ImageNet特征也能提供边缘纹理基元能让损失函数下降得更平稳。6.2 测试时增强TTA推理时对输入做水平翻转、垂直翻转和90度旋转医学图像轴向旋转90度要谨慎得到多组概率图取平均后预测。下面是一个简单的TTA封装def tta_predict(model, volume): probs [] with torch.no_grad(): for flip in [False, True]: v torch.flip(volume, dims[-1]) if flip else volume p torch.softmax(model(v), dim1) if flip: p torch.flip(p, dims[-1]) probs.append(p) return torch.stack(probs).mean(0)6.3 用多尺度预测修正边缘对输入分别用原始大小和0.5倍分辨率跑一遍把低分辨率结果上采样回原始大小和原始分辨率预测平均。低分辨率感受野更大对大面积漏检有帮助原始分辨率保留细节。这个技巧和TTA可以叠加在我遇到的肝脏MRI数据上Dice能提升1-2个点。6.4 检查方向如果预测mask整体偏移如果分割结果比金标准小一圈或者存在一致偏移先检查预处理里的裁剪和重采样是否改变了体素间距再检查DataLoader shuffle时是否把图像和标签配对错了。另一个常见的坑是训练时用了padding“same”而推理时没有导致feature map尺寸不一致输出被截断。调试时可以打印推理输出的logits尺寸和输入尺寸确保完全一致。最后再说一句不要盲目追求3D模型。先用2.5D方案跑通整个流程再逐步换模型和调参这样出结果最快也最容易定位问题。本文还有配套的精品资源点击获取
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。