ARTICLE DETAIL

资讯详情

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

UNet医学图像分割:多类别分割从数据准备到训练避坑指南

UNet医学图像分割:多类别分割从数据准备到训练避坑指南 简介面向医学图像分割、语义分割与多类别分割任务的U-Net实现代码包源自CSDN作者qq_44886601适合深度学习者、科研人员及医工交叉领域开发者快速复现经典分割模型。U-Net采用对称的收缩-扩展路径借助跳跃连接融合上下文信息与高分辨率特征对医学影像中边界模糊区域具有良好分割表现。压缩包共31个文件以Python源码.py与编译缓存.pyc为主另含依赖清单、配置xml及readme说明整体仅16KB轻量便携。代码模块覆盖模型定义、数据加载、数据增强、训练预测及混淆矩阵评估等完整流程可直接运行并支持按需调整既便于理解U-Net架构设计也可迁移到其他语义分割或多类别分割应用。已有466人浏览学习适合用于项目实战、课程设计或论文复现能帮助读者从数据处理到模型评估建立完整的分割流水线。1. unet 医学图像分割为什么它仍是多类别分割任务的默认基线做过医学图像分割的人大致都有同感刚拿到带标注的 CT 或 MRI 数据时第一反应不是去追最新的 Transformer 结构而是先把手里的 unet 代码跑起来出一个 baseline。unet 医学图像分割之所以流行这么多年不是因为结构花哨而是因为在标注样本普遍偏少、图像噪声偏偏又大的医学场景里它用一条很朴素的思路——下采样抓上下文、上采样恢复细节、跳跃连接保边缘——就把分割任务的大头拿下了。本文要聊的是怎么把一份 unet 多类别分割代码从头到尾真正跑通从数据读取到损失设计再到避坑覆盖语义分割和多类别分割在医学场景下的完整落地路径。适合正在做医学图像相关课题的研究生、刚转算法岗的工程师以及任何准备拿 unet 训练自己数据集但还没理清细节的人。2. 拆解 unet 结构编码器、跳跃连接与输出通道数的选择2.1 编码器与解码器为什么医学图像吃这套结构unet 的网络结构一句话就能说清左边一条编码器路径逐级下采样缩小特征图尺寸并增加通道数右边一条解码器路径逐级上采样恢复空间分辨率中间用跳跃连接把同尺度的编码器特征直接拼到解码器上。这个设计对医学图像非常契合因为医学影像里要分割的目标往往边缘模糊、和周围组织灰度接近比如肝脏肿瘤在 CT 里就可能只差几个 HU 值。编码器一路压缩模型被迫学会“看全局”但纯压缩会丢失边缘细节所以跳跃连接又把浅层的高分辨率特征原样送到解码器让模型在恢复轮廓时有足够的信息可参考。我在实际写代码时最常被问到的一个问题是到底要不要用预训练 backbone如果是自然图像数据集用 torchvision 里预训练好的 resnet 当编码器通常能涨点因为 ImageNet 上学到的纹理特征可以迁移。但医学图像是灰度图为主模态差异太大预训练权重带来的增益不稳定。我自己的经验是公开的医学分割数据集上从零训练一个基础版 unet 往往已经够了如果样本特别少优先做数据增强比换预训练模型更划算。2.2 跳跃连接能不能去掉精度与显存的取舍跳跃连接是 unet 区别于普通自编码器的关键不建议去掉。但要知道它的代价每一次跳跃连接都要把编码器的特征图和解码器当前的特征图在通道维度做拼接这直接推高了显存占用。比如最底层解码器的输入瓶颈层上采样后是 512 通道拼接上编码器第 4 层的 512 通道特征变成 1024 通道再进卷积块这一层几乎是全网络显存峰值所在。图像尺寸是 512×512 时这个拼接的显存开销还能接受一旦换成 2D 切片训练但图像分辨率到了 1024 或更大显存就会吃紧。常见做法是在保证 batch size 不为 1 的情况下把输入缩到 256×256 或 512×512既能拿到足够多的空间细节又不至于让显存成为瓶颈。拼接方式本身也可以做微调有人改成相加add代替拼接concat来省显存但这会轻微掉精度我一般不轻易改。2.3 多类别分割的输出通道设计别把类别数搞混unet 多类别分割和普通二分类分割的唯一结构性差异在最后一层。二分类通常输出 1 个通道接 sigmoid多类别分割输出通道数等于类别数接 softmax。比如分割肝脏、肾脏、脾脏三个器官加上背景就是 4 类输出层就是 4 个通道。这里有个很容易让新手翻车的细节你的标签掩码如果是 0、1、2、3 这样的整数索引那么标签 tenser 要用 long 类型损失函数用 CrossEntropyLoss如果你的掩码是 one-hot 编码那就要在损失函数里先做 argmax 转回整数索引再算交叉熵。两种写法都能通但千万别一种标签配了另一种损失训练 loss 会一路飘着降不下来而且是那种怎么看都有问题又不报错的玄学状态——模型没炸但指标就是不动。3. 医学图像分割数据集准备从 NIfTI 切片到多类别掩码3.1 读 NIfTI 并切成 2D 切片先搞清楚坐标轴医学影像最常见的存储格式是 NIfTI.nii / .nii.gz。它和普通 PNG 的最大区别是多了一个 world coordinate 信息但训练 unet 时我们一般只关心体素值。用 SimpleITK 读文件并转成 numpy 数组最需要注意的是 GetArrayFromImage 返回的维度顺序是 (z, y, x)也就是说第 0 维是层数第 1、2 维才是图像平面。import numpy as np import SimpleITK as sitk def load_nii_slice(image_path, label_path, slice_idx): # 读入 nii 文件返回 3D 数组 image sitk.GetArrayFromImage(sitk.ReadImage(image_path)) # (z, y, x) label sitk.GetArrayFromImage(sitk.ReadImage(label_path)) # (z, y, x) return image[slice_idx], label[slice_idx]这段代码的逻辑很直白把 nii 的 3D 体数据沿 z 轴切成一张张 2D 切片取第 slice_idx 张作为训练样本。切片方向怎么选取决于你的任务——肝脏 CT 通常沿横断面切也就是第 0 维直接取但如果是血管类数据可能要沿冠状面或矢状面切才能看到连续结构。我一般会在开始训练前把三个方向的切片各随机看几张确定哪个方向的目标形态最稳定再决定主切方向。3.2 掩码标签处理整数索引标签和 one-hot 标签的换算读出来的标签掩码一般有两种长相一种是 0、1、2、3 的整数数组每个数值代表一个类别另一种是向量图导出时把每一类单独存成一个 255 的图层比如 mask_class1.png 里白色是肝脏、mask_class2.png 里白色是肾脏这类掩码在医学公开数据集里很常见需要先合成再转索引。import cv2 import numpy as np def merge_masks_to_indices(mask_list, output_dtypenp.int64): # mask_list: 每张图同尺寸1 表示前景0 表示背景 merged np.zeros(mask_list[0].shape, dtypeoutput_dtype) for idx, mask in enumerate(mask_list): # 逐类别填索引从 1 开始0 留给背景 merged[mask 0] idx 1 return merged这里必须留意的坑是类别索引从 1 而不是从 0 开始0 永远留给背景。如果某张图中一个类别的掩码是空的即所有像素都为 0那 merged 里这一类的索引就不会出现计算损失和指标时就要跳过空类别否则会在评价函数里除零。后续训练代码里CrossEntropyLoss 的类别权重也建议按类别像素占比倒过来给这个放到第 4 章细说。3.3 数据增强旋转、翻转与弹性形变的参数设计医学数据量小增强是提高泛化能力最廉价的手段。我常用的增强组合是随机旋转 ±15 度、水平翻转、随机亮度对比度以及轻微的弹性形变。用 albumentations 库写起来最顺手因为它能对图像和掩码同步做同样的几何变换。import albumentations as A train_transform A.Compose([ A.RandomRotate90(p0.5), A.HorizontalFlip(p0.5), A.RandomBrightnessContrast(brightness_limit0.1, contrast_limit0.1, p0.5), A.ElasticTransform(alpha1.0, sigma10.0, alpha_affine10.0, p0.3), A.Resize(512, 512), ])参数说明写在代码后面的习惯我保持了很久。alpha 和 sigma 控制弹性形变的幅度医学器官的形变通常是整体推移而不是局部扭曲因此 alpha 不宜太大我一般 alpha1.0、sigma10.0 起调。Resize 到 512×512 是最稳当的选择再大显存容易爆再小会丢失小病灶细节。注意 RandomBrightnessContrast 在增强列表里的顺序放在几何变换之后、Resize 之后都行但不要先 Resize 再做随机旋转否则会引入黑色边框噪声。4. 训练 unet 多类别分割模型损失函数设计与核心代码4.1 损失函数选型Dice、交叉熵还是组合损失多类别分割的损失函数选择基本决定了训练的最后高度。单独用交叉熵小器官类别容易因为像素占比低而被模型直接忽略单独用 Dice Loss训练前期梯度噪声大收敛不稳定。我自己的常规操作是交叉熵加 Dice 的加权组合交叉熵负责稳定收敛Dice 负责拉高目标区域的 IoU。这里补充一点稍微进阶的选择如果你觉得背景占比过大比如一张 512×512 的切片里器官只占不到 5%那 Dice 部分虽然已经天然缓解类别不平衡但还是有限。可以改成 Tversky Loss在 Dice 基础上多引入两个超参数 alpha 和 beta 来惩罚假阳性和假阴性。不过这是后话对大多数医学数据集Dice 加 CE 的组合已经够用不建议一上来就调一大堆超参。4.2 多类别训练代码从模型定义到训练主循环把 unet 模型、混合损失和训练主循环串起来是整个过程里最核心的部分。模型代码我用的是最经典的双卷积加拼接写法没有额外加注意力模块先把 baseline 跑通再说。import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue) ) def forward(self, x): return self.conv(x) class UNet(nn.Module): def __init__(self, in_channels1, num_classes4): super().__init__() self.enc1 DoubleConv(in_channels, 64) self.enc2 DoubleConv(64, 128) self.enc3 DoubleConv(128, 256) self.enc4 DoubleConv(256, 512) self.pool nn.MaxPool2d(2) self.bottleneck DoubleConv(512, 1024) self.up4 nn.ConvTranspose2d(1024, 512, 2, stride2) self.dec4 DoubleConv(1024, 512) self.up3 nn.ConvTranspose2d(512, 256, 2, stride2) self.dec3 DoubleConv(512, 256) self.up2 nn.ConvTranspose2d(256, 128, 2, stride2) self.dec2 DoubleConv(256, 128) self.up1 nn.ConvTranspose2d(128, 64, 2, stride2) self.dec1 DoubleConv(128, 64) self.out_conv nn.Conv2d(64, num_classes, 1) def forward(self, x): e1 self.enc1(x) e2 self.enc2(self.pool(e1)) e3 self.enc3(self.pool(e2)) e4 self.enc4(self.pool(e3)) b self.bottleneck(self.pool(e4)) d4 self.dec4(torch.cat([self.up4(b), e4], dim1)) d3 self.dec3(torch.cat([self.up3(d4), e3], dim1)) d2 self.dec2(torch.cat([self.up2(d3), e2], dim1)) d1 self.dec1(torch.cat([self.up1(d2), e1], dim1)) return self.out_conv(d1)in_channels 要对应输入图像的通道数CT 单序列是 1MRI 如果有 T1、T2 两个序列叠加就是 2RGB 染色病理图是 3。num_classes 是包括背景在内的总类别数务必不要只填器官个数否则输出通道少一个会把整个类别空间的 softmax 算错。接下来是损失函数的实现。交叉熵和 Dice 组合的关键是记得交叉熵输入是原始 logits而 Dice 需要先做 softmax 把 logits 转成概率再计算def dice_loss(pred, target, smooth1.0): pred torch.softmax(pred, dim1) # (B, C, H, W) target_onehot torch.nn.functional.one_hot( target, num_classespred.shape[1] ).permute(0, 3, 1, 2).float() # (B, C, H, W) intersection (pred * target_onehot).sum(dim(2, 3)) union pred.sum(dim(2, 3)) target_onehot.sum(dim(2, 3)) dice (2 * intersection smooth) / (union smooth) return 1 - dice.mean() def mixed_loss(pred, target): ce nn.CrossEntropyLoss()(pred, target) dl dice_loss(pred, target) return ce dl这里 target 是整数索引标签形状是 (B, H, W)。dice_loss 里 one_hot 在类别维度展开让我对每个类别单独算 Dice 再取平均。smooth 建议给 1.0防止某个类别在切片上完全没有出现时分子分母全零。初次训练如果觉得模型偏向忽略小目标可以把混合损失改成 ce 2 * dl也就是把 Dice 的权重加大让模型更关注区域重叠度。最后是训练主循环。医学分割图像的尺寸大训练代码里一定要加梯度裁剪不然遇到异常的增强样本很容易梯度爆炸train_loader DataLoader(train_dataset, batch_size8, shuffleTrue, num_workers4, drop_lastTrue) optimizer torch.optim.Adam(model.parameters(), lr1e-3) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size30, gamma0.5) model model.to(device) for epoch in range(80): model.train() total_loss 0.0 for image, mask in train_loader: image image.to(device).float() mask mask.to(device).long() pred model(image) loss mixed_loss(pred, mask) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 12.0) optimizer.step() total_loss loss.item() scheduler.step() val_dice evaluate(model, val_loader, device) print(fepoch {epoch:3d} | loss {total_loss/len(train_loader):.4f} | val dice {val_dice:.4f})学习率 1e-3 配合 Adam 是医学分割训练的常见起点如果你的 batch size 很小比如只有 4建议把学习率降到 5e-4否则 BatchNorm 统计量不稳定容易震荡。StepLR 每 30 轮降一半我一般要看到验证集指标开始平台期才愿意等它触发所以也经常直接用 ReduceLROnPlateau 看验证集指标动态降。drop_lastTrue 是小心机避免最后一个 batch 样本数过少导致的 BatchNorm 抖动。4.3 多类别评估指标Dice 与 mIoU 的适用场景差异验证阶段常见有两个指标打架mIoU 和 Dice。很多遥感或自然图像的语义分割任务习惯报 mIoU因为 IoU 对背景占比大的情况更敏感但医学分割论文几乎都报 Dice因为它更贴近病灶区域的真实重叠程度。两者数值上并不是简单的倍数关系——Dice 通常比同条件下的 mIoU 高 5 到 10 个百分点所以在对比别人成果时先确认对方报的是什么指标再比较这是最容易出理解偏差的地方。def compute_dice(pred_mask, gt_mask, num_classes): dice_list [] for cls in range(1, num_classes): # 跳过背景 p (pred_mask cls) t (gt_mask cls) inter (p t).sum() total p.sum() t.sum() dice 2 * inter / total if total 0 else 1.0 dice_list.append(dice) return np.mean(dice_list)这里计算每个类别的 Dice 后再平均叫 macro Dice会给小类别更高的权重符合医学分割的关注重点。如果某个类在验证集里完全没有出现过直接给 1.0 而不是 0.0因为这一类的预测结果等于是“处处为空、预测为空”属于完全正确的空预测。如果给 0会把整个验证集指标拖得惨不忍睹而且这个偏差在类别很少时会被放大。5. unet 分割常见问题排查五条血泪踩坑记录5.1 现象一Dice Loss 训着训着变成 NaN训练到中间某个 epochloss 突然变成 NaN然后一路 NaN 到底。最初我以为是学习率太大降下去还是复现后来发现是标签里有某个类别的掩码整张为空one_hot 后目标概率为 0而预测概率又非常接近 0交叉熵计算时 log(0) 直接产出无穷大。解决方法是给交叉熵加 label_smoothing或者把空类别的索引从训练集里剔除另一个常见来源是 BatchNorm 在 batch size 为 1 时统计量退化遇到这种就检查一下 drop_last 和 batch size 是否过小。5.2 现象二多类别分割里小器官永远分不出来胰腺、小型淋巴结这类目标在 CT 切片上往往只占几百个像素模型在竞争中天然偏向大器官。先别急着换损失函数优先检查预处理里有没有做窗宽窗位调整——腹部 CT 如果不调窗肝胰脾在灰度上几乎是融在一起的模型再强也分不开。做完图像前处理还没改善再考虑加入类别权重在 CrossEntropyLoss 的 weight 参数里把像素占比小的类别权重调大同时把输入分辨率从 256 提到 512小器官的边界特征会明显好找很多。5.3 现象三验证集 Dice 不低可视化却全是碎点这种情况经常出现在分割精度高但预测图噪声多的模型里Dice 是像素级指标一个目标上多个小块碎点总面积占比不高指标看起来过得去但医生根本无法临床使用。根因是训练时没有加连通域相关的约束模型在决策边界处抖动剧烈。工业界的常规操作是在推理后处理里做连通域分析只保留面积最大的连通域或者是按类别设定最小面积阈值碎点直接删掉。这一步不吃任何训练资源却能显著拉开“指标好看”和“结果能用”之间的距离。5.4 现象四显存不够导致喂不了大图512×512 的输入batch size 8一张普通消费级 12G 显存卡勉强能跑换成 1024×1024 直接爆显存这是医学图像分辨率带来的常见卡点。两个方案并行最好一个是把 batch size 降到 2 甚至 1配合梯度累积模拟更大 batch另一个是把整图切成 patch 输入比如切成 256×256 的 patch训练时每个 patch 只包含局部器官推理时用滑动窗口拼接。切 patch 后要注意标签掩码里类别分布会变某些 patch 可能全是背景这种样本要在采样时过滤掉不然会拉低训练效率。5.5 现象五模型不收敛且预测结果全部是背景训练 loss 在一两个 epoch 后快速下降验证集 Dice 却是 0可视化后模型输出全黑。这类问题九成是数据预处理阶段标题就错了CT 原始数值范围是 -1000 到 3000 左右如果没做归一化直接喂进网络模型前几层的梯度会来回震荡反过来如果只做了除以 255窗口又没落在软组织范围上器官和背景的对比度也出不来。规范做法是先用窗宽窗位把数值裁到目标组织范围比如腹部软组织是 -125 到 225再线性映射到 0 到 1最后如果还不行把标签拖到画布上逐类检查有没有类别错位。6. 让分割结果更可用后处理技巧与两个值得尝试的改进方向训练结束只是第一步真正交付给下游使用的分割结果还需要处理后才能看。我习惯在推理之后做两件事先用形状学闭运算把细小的断裂连接起来再用连通域过滤去掉孤立噪声点。闭运算的核大小不要超过 5×5太大的核会把两个本不相邻的器官粘连在一起连通域过滤则按类别单独设阈值小器官阈值设低一点大器官阈值设高一点。推理输出 (B, C, H, W) - argmax 得到 (B, H, W) 整数掩码 - 对每个类别做连通域分析 - 删除像素数小于阈值 category_min_size 的连通域 - 得到最终掩码这一套处理可以放在验证循环里跑一遍对比后处理前后的 Dice 变化通常能涨 0.5 到 1.5 个点而且涨的是真正有用的精度。如果 baseline 稳定跑通还想进一步涨点两个方向值得花时间。第一个是 unet 本身的修改常见思路是在跳跃连接处加入注意力机制让解码器自动选择更重要的编码器特征层最经典的变体是 Attention U-Net第二个是把编码器换成预训练分类网络比如 ResNet34 做编码器、保留 unet 的解码器结构这类改进在中等规模数据集上通常有明显收益。我不太建议一开始就上 TransUNet 或 Swin-Unet显存开销大、收敛慢先把手写 unet 的每个环节跑扎实更划算。我自己的习惯是每跑一次实验就把训练集 Dice、验证集 Dice、后处理后 Dice 和可视化图放到一张表里对照数字再好看图上一片碎点也是白搭。这个习惯帮我拦下了不少自我感觉良好的翻车实验。希望帮到你。本文还有配套的精品资源点击获取
返回列表