基于PyTorch的UNet多类别医学图像分割实践与避坑指南
2026/9/24 20:17:58 网站建设 项目流程

简介:一份面向医学图像分割与语义分割任务的U-Net工程代码,适合深度学习者、医学影像算法入门者快速开展实验。资源实现了经典的U-Net网络结构,覆盖单/多类别分割场景,可应用于病灶定位、组织器官轮廓提取等任务。压缩包共31个文件,主体为8个Python脚本,涵盖数据集构建、模型定义、训练、预测及混淆矩阵评估等完整流程;另含14个pyc编译文件和工程配置文件,便于复用与调试,整体仅16KB,轻量易部署。目前已有466人学习使用,说明其实用性受到一定认可。通过阅读README和源码,可掌握U-Net的跳跃连接设计与训练细节,也能直接替换自己的医学图像数据集进行迁移训练。

1. 多类别医学图像分割的起点:UNet 代码到底在解决什么

一个很常见的场景:脑肿瘤 MRI 用 UNet 做二分类分割,跑几天能出不错的结果;一旦换成多类别分割,要么边缘糊成一团,要么细小的血管直接消失,甚至验证集指标不错、可视化却发现某些类别完全没被预测出来。UNet 能在医学图像分割里成为默认选择,靠的是跳跃连接把编码器的高层语义和解码器的空间细节直接拼起来,对小器官、模糊边界有很高的容错性。下面把数据格式、归一化、多类别输出、损失函数与训练推理的完整代码链路理一遍,适合刚拿到公开数据集、准备把单类别语义分割扩展成多类别任务的开发者。

2. 把医学图像转成 UNet 能吃的张量:格式、归一化与多类别合并

2.1 三类常见数据形式:nii/nii.gz、png 掩码与 numpy 数组

医学图像分割的数据和自然图像差别很大。最常见的是 nii/nii.gz 体数据,通常由 DICOM 序列转换而来,内部包含完整的 spacing、origin、direction 信息;其次是同步保存的 mask 文件,可能是一整个 npy 数组,也可能是一组按切片保存的 png。开始写代码前,先把所有输入统一成一个约定:图像读进来转成 float32,标签转成 uint8,类别从 0 开始连续编号。

我一般会先搭一个简单的目录结构,避免后面在 DataLoader 里反复改路径:

data/ images/ case_01.nii.gz case_02.nii.gz masks/ case_01.npy case_02.npy

读取 nii 时建议直接用 SimpleITK,它能保留体数据的空间元信息,训练时用不到,但推理结束要写回 nii 给医生看时,原图的 spacing 和 origin 必须原样带回去,否则在专业阅片软件里会错位。读取代码长这样:

import numpy as np import SimpleITK as sitk def load_volume(path): img = sitk.ReadImage(path) # 一次读入完整 nii,保留元信息 arr = sitk.GetArrayFromImage(img) # 转成 numpy,形状为 (depth, H, W) return arr.astype(np.float32), img volume, img_meta = load_volume("data/images/case_01.nii.gz")

这里有两个参数层面的关键点:SimpleITK 读出来的数组是 (depth, H, W),和常见 2D 数据集的 (H, W) 不同,进 UNet 前要把单张切片取出来;img_meta 对象不能丢,后面推理完要拿它恢复方向。另一点是 dtype,float32 足够训练,没必要用 float64 占显存。

如果数据集直接给了 png 掩码,读取时注意通道数问题。很多医学 png 虽然看起来是灰的,但可能保存成了三通道。统一用 cv2.IMREAD_GRAYSCALE 强制读成单通道,比读进来再 squeeze 更稳,因为有些三通道灰度图 squeeze 后形状会乱。

2.2 归一化策略与多类别标签合并:两个容易做错的细节

CT 与 MRI 的归一化逻辑完全不同。CT 的像素值本质是亨氏单位(HU),范围通常在 -1024 到 3071 之间,直接除以 255 或做 min-max 归一化都会让软组织细节被压没,因为肺、脂肪、骨头的 HU 差异很大,但肿瘤往往只落在很窄的窗口里。常规做法是先做窗宽窗位截断,再归一化:

def normalize_ct(volume, window_min=-200, window_max=200): volume = np.clip(volume, window_min, window_max) volume = (volume - window_min) / (window_max - window_min) return volume.astype(np.float32)

窗口参数不是固定的,肝部病灶和肺部结节的最佳窗口差别很大,需要配合具体标注范围做调整。MRI 没有统一的物理单位,不同序列之间的亮度含义完全不同,常规做法是 percentile 裁剪或 z-score 标准化;直接用 ImageNet 上统计的 mean/std 不适用,因为那是三通道自然图像的统计量。

多类别标签合并是另一个高频踩坑点。公开数据集里 mask 的保存方式五花八门:有的按类别分文件,有一个是 tumor.png、另一个是 edema.png;有的单张 png 里用 0、128、255 表示不同目标;还有一个 npy 里直接是整数数组但类别不是从 0 开始的。训练前必须统一成从 0 开始的连续标签图:

label = np.zeros((H, W), dtype=np.uint8) label[edema_mask > 0] = 1 label[tumor_core_mask > 0] = 2 # 如果不同类别在空间上重叠,自己定优先级,后赋值会覆盖先赋值 print(np.unique(label)) # 必须看到 [0 1 2] 这种连续结果

赋值顺序就是优先级,重叠区域谁后赋值谁生效。如果直接从 png 读,像素值 128 和 255 不能直接用,要先把它们重映射成 1 和 2。训练前把这个 np.unique 的检查写成脚本,能省掉后面至少两轮排查时间。

3. 用 PyTorch 搭 UNet:从编码器到多类别输出的完整代码

3.1 核心结构:双卷积模块、下采样与跳跃连接

UNet 的骨架可以拆成三块:编码器负责逐步下采样提取语义,解码器负责逐步恢复分辨率,跳跃连接把编码器每一层特征拼到解码器的对应层。跳跃连接是 UNet 的题眼,它让解码器在恢复细节时既能看到高层语义,又能直接引用最底层的边缘纹理,对小器官分割特别关键。

先写最基础的 DoubleConv 模块。每一步都做两次卷积,而不是一次,是为了让每个尺度上的特征表达更充分。卷积用 padding=1 保持特征图尺寸不缩水,下采样交给单独的 MaxPool2d:

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, kernel_size=3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) def forward(self, x): return self.conv(x)

kernel_size=3、padding=1 的组合保证输入输出尺寸一致,后面拼接时不用做尺寸对齐;BatchNorm 放在卷积之后、激活之前是 PyTorch 里最稳的写法。两轮卷积的作用不是叠更深,而是让每个尺度都有足够的感受野去捕捉局部纹理。

完整网络按常见配置写:初始通道数 64,每下采样一次通道翻倍,到最底层变成 512。通道数和数据集大小强相关,小数据集起步 32 就够了,通道翻倍太狠反而容易在小样本上学不到有效特征。

class UNet(nn.Module): def __init__(self, in_ch=1, num_classes=3): super().__init__() self.inc = DoubleConv(in_ch, 64) self.down1 = DoubleConv(64, 128) self.down2 = DoubleConv(128, 256) self.down3 = DoubleConv(256, 512) self.pool = nn.MaxPool2d(2) self.up3 = nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2) self.dec3 = DoubleConv(512, 256) self.up2 = nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2) self.dec2 = DoubleConv(256, 128) self.up1 = nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2) self.dec1 = DoubleConv(128, 64) self.out_conv = nn.Conv2d(64, num_classes, kernel_size=1) def forward(self, x): x1 = self.inc(x) x2 = self.down1(self.pool(x1)) x3 = self.down2(self.pool(x2)) x4 = self.down3(self.pool(x3)) x = self.up3(x4) x = torch.cat([x, x3], dim=1) # 跳跃连接:通道维拼接 x = self.dec3(x) x = self.up2(x) x = torch.cat([x, x2], dim=1) x = self.dec2(x) x = self.up1(x) x = torch.cat([x, x1], dim=1) x = self.dec1(x) return self.out_conv(x)

上采样用 ConvTranspose2d,kernel_size=2、stride=2 是反卷积里最基础的尺寸翻倍配置。torch.cat 在 dim=1 上拼通道,所以拼接后通道数翻倍,对应 DoubleConv 的输入通道要写成 512、256、128。如果输入图像尺寸不是 16 的整数倍,下采样到最底层时会出现 ceil 和 floor 的差异,拼接时尺寸差 1 像素会直接报错,医学数据最稳的方法是训练前统一把图像 resize 或 crop 成 256x256、512x512 这类 2 的幂尺寸。

3.2 改造成多类别分割:输出通道、激活函数与损失函数怎么对接

把 UNet 从二分类改成多类别,结构上只需要改两个地方:in_ch 改成输入图像的通道数,灰度图是 1,RGB 自然图是 3;num_classes 改成标注里的总类别数。最后一层输出 num_classes 个通道,每个通道对应一个类别的预测得分,也就是 logits。

这里最容易被带偏的是激活函数的选择。多类别分割里如果每个像素只属于一个类别,最后一层就不能加 sigmoid,而要配合 CrossEntropyLoss 内部的 softmax 使用;只有任务本身允许多个类别同时出现在同一个像素上,才用 sigmoid 配合 BCEWithLogitsLoss。医学图像里大多数逐像素标注是互斥的,肿瘤不可能同时是背景,所以常规做法是写 num_classes 输出、不接激活:

# 前向输出 logits,shape 为 (B, num_classes, H, W) logits = model(x) # 训练时直接交给 CrossEntropyLoss loss = nn.CrossEntropyLoss()(logits, target) # 推理时才能做 argmax pred = torch.argmax(torch.softmax(logits, dim=1), dim=1)

训练阶段的 CrossEntropyLoss 对 target 有隐式要求:target 的数据类型必须是 long,且数值范围必须在 [0, num_classes-1] 之间。很多人在这里翻车,比如标签里混了 255 或 3 类以外的值,损失函数不会直接报错,但会把这些错值当 ignore_index 处理,导致类别永远学不出来。目标检测也同理,语义分割只不过把边界框换成了逐像素掩码,输出矩阵的最后一维代表类别概率。

至于常见的 UNet 结构改进,比如把 DoubleConv 换成 ResBlock、在跳跃连接上加 attention gate,本质上都是在这个骨架上做特征筛选。新手先跑通基础版,验证数据、损失和指标没毛病,再谈改进,否则改结构和调数据问题混在一起,出了问题根本定位不到原因。

4. 训练配置与验证指标:把 Dice、IoU 的坑留在进入实验之前

4.1 损失函数选型:CrossEntropy、Dice Loss 与组合损失

多类别分割的损失函数选择直接决定训练能否收敛。CrossEntropyLoss 对每个像素独立计算,实现简单、数值稳定,但在类别极不均衡时会被占比大的背景类主导。医学图像里肿瘤区域往往只占全图的 1% 到 5%,单用交叉熵时模型很容易学到背景而忽略病灶。

Dice Loss 是医学分割里更常用的一类损失,它直接优化分割结果与金标准之间的区域重叠,对小目标更友好。多类别下的标准实现是按每个类别分别算 Dice 再取平均:

def dice_loss_multiclass(logits, target, class_weights=None): probs = torch.softmax(logits, dim=1) n_classes = probs.shape[1] dice_sum = 0.0 total_weight = 0.0 for c in range(n_classes): p = probs[:, c] # 当前类别的预测概率 t = (target == c).float() # 当前类别的真实掩码 inter = (p * t).sum() union = p.sum() + t.sum() + 1e-6 dice_c = 2 * inter / union if class_weights is not None: dice_c = dice_c * class_weights[c] total_weight += class_weights[c] dice_sum += dice_c return 1 - dice_sum / (total_weight if class_weights is not None else n_classes)

class_weights 参数可用于放大稀有类别的贡献,比如类别 2 只占全图 2%,权重设成 2.0 能让模型在训练早期重点关注它。epsilon 取 1e-6 是为了防止某类在某个 batch 完全没出现时出现除零。另外,背景类是否参与 Dice 计算要谨慎:如果背景占比压倒性大,背景的 Dice 会拖着整体数值虚高,验证时看起来不错,实际病灶区域一塌糊涂。

组合损失是更稳的默认做法。我会把 CrossEntropy 和 Dice Loss 按权重加起来,常见比例是 0.5:0.5;当类别严重不均衡时调整到 0.3:0.7,让 Dice 主导。这样既保留交叉熵的梯度稳定性,又让模型直接对准区域重叠目标。实现时需要注意两个 loss 的取值范围不同,CE 的典型数值比 Dice Loss 大一个量级,需要各自加一个可学习或手动调的权重系数,否则等价于把 CE 忽略掉。

4.2 batch size、学习率与验证指标:参数配置表和可复现训练循环

训练配置里最影响结果的三件事:学习率、batch size、验证指标的计算方式。医学图像单张尺寸大,显存经常卡在 batch size 上。一张 512x512 的灰度图在 batch size 8 时,UNet 大约要占 8-10GB 显存,很多卡只能勉强跑 batch size 2。但 batch size 太小带来两个连锁问题:BatchNorm 的统计量不稳定,模型每步更新方向抖动剧烈。常规思路是先用较小的输入尺寸,例如 256x256 crop 把 batch size 拉大到 8 以上,而不是在 512x512 硬撑 batch size 2 然后怀疑模型不收敛。

损失函数后面的反向传播逻辑是我经常保存的标准模板,整个循环要保证验证阶段不更新梯度,同时把预测和标签从 CUDA 显存里拿回来算指标:

criterion = nn.CrossEntropyLoss() optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode="max", patience=8, factor=0.5) for epoch in range(epochs): model.train() train_loss = 0.0 for images, masks in train_loader: images = images.float().to(device) masks = masks.long().to(device) pred = model(images) loss = criterion(pred, masks) optimizer.zero_grad() loss.backward() optimizer.step() train_loss += loss.item() model.eval() val_dice = compute_mean_dice(model, val_loader, device) scheduler.step(val_dice)

学习率设置在医学图像分割里常见做法是 1e-4 起步,这是 Adam 家族在这个任务上的经验值。如果 loss 在最初几个 epoch 完全不动,优先检查标签映射是否连续、输入归一化是否正常,而不是急着调学习率。ReduceLROnPlateau 的 mode 要配 max,因为监控的目标是 Dice 而非 loss,Dice 越大越好;这也是一个常见的配置误区。

验证指标不能只看整体 Dice,写一个按类别分别计算的函数:

def compute_mean_dice(model, loader, device, num_classes=3): model.eval() class_dice = [0.0] * num_classes samples = [0] * num_classes with torch.no_grad(): for images, masks in loader: pred = model(images.to(device)) pred = torch.argmax(pred, dim=1).cpu() for c in range(num_classes): p = (pred.numpy() == c) t = (masks.numpy() == c) if t.sum() > 0: inter = (p & t).sum() union = p.sum() + t.sum() class_dice[c] += 2 * inter / union samples[c] += 1 return np.mean([class_dice[c] / max(samples[c], 1) for c in range(num_classes)])

类别 c 在某个 batch 里完全不存在时,不能直接记一个 0 去平均,否则验证集大的时候指标被无意义地拉低。更现实的做法是只在真实掩码包含该类别的样本上累加,最后用计数做分母。打印指标时,我习惯把每个类别的 Dice 单独输出一遍,只有 mean Dice 一个数很容易掩盖个别类别根本没学会的问题。

数据增强在医学分割里也值得多说一句。翻转、旋转、弹性形变是最常用的三类增强,但增强必须作用在图像和掩码上保持同步,否则等于在给模型喂标签错误的数据。albumentations 库对这类需求支持比较好,它把 image 和 mask 放在同一个 transform pipeline 里。对于小数据集,弹性形变带来的提升往往比换模型结构更明显,这是个不需要额外计算成本的白嫖技巧。

5. 多类别分割避坑记录:五个让我白跑实验的问题与排查

5.1 坑一:灰度图读进来却是三通道,输入尺寸对不上

现象:训练 loss 能正常下降,一到验证阶段喂验证集就报 tensor shape 不匹配,提示输入维度是 (B, 3, H, W) 但模型第一层是 Conv2d(1, …)。

原因:读取 png 掩码时用了 cv2.imread 默认的 IMREAD_COLOR 模式,灰度图被复制成三通道。三通道虽然都是同一份灰度数据,但模型第一层的 in_ch=1 直接拒绝。

解决:统一用 cv2.imread(path, cv2.IMREAD_GRAYSCALE),并把这一约束写进数据读取函数开头。输入图像也建议显示打印一次 shape,把这个检查留在吐槽模型玄学之前。

5.2 坑二:batch size=1 时 BatchNorm 在验证集上疯狂抖动

现象:训练 loss 稳步下降,每轮验证的 Dice 忽高忽低,同一份模型不修改代码重跑一遍,结果又不一样。

原因:BatchNorm2d 在 batch size=1 时,单样本的均值和方差就是它自身,归一化退化成线性缩放,每步都在震荡,模型学到了不稳定的分布。

解决:尽量把 batch size 提到 4 以上。显存不够时先缩小输入尺寸或裁剪 patch,而不是硬扛大图;如果任务确实要求单样本推理,把 BatchNorm2d 换成 InstanceNorm2d,后者对 batch size 不敏感,医学小样本分割里替换后通常会稳很多。

5.3 坑三:标签类别不连续,CrossEntropyLoss 静默背锅

现象:训练 loss 降到一个很低的值,但可视化预测结果时发现某几个类别从来没被预测出来,Dice 再高也没用。

原因:公开数据集的标签文件里混入了 255 或缺失的类别编号。CrossEntropyLoss 遇到 target=255 时默认忽略该像素,类别 2 如果恰好缺失,模型从未在这个类别上收到梯度,自然学不会。

解决:数据预处理阶段强制 assert,类别必须严格等于 [0, 1, 2]:

assert set(np.unique(label)).issubset(set(range(num_classes))), \ f"标签异常: {np.unique(label)}, 期望类别 0-{num_classes-1}"

5.4 坑四:多类别分割误用 sigmoid 加 BCE

现象:每个像素能且只能属于一个类别,但用 sigmoid + BCEWithLogitsLoss 训练后,预测图里经常出现多个类别同时为 1,而且边界区域互相重叠打架。

原因:BCE 对每个类别独立判断,等价于多标签任务,它不强制类别之间互斥。多类别分割的语义是多分类,必须用 softmax 输出类别分布,用 CrossEntropyLoss。

解决:检查任务标注是否允许多标签。绝大多数语义分割数据集是单标签的,直接统一用 CrossEntropyLoss。只有在同一个体素可以同时属于血管和病变这类重叠分割需求时,才保留 sigmoid + BCE。

5.5 坑五:数据增强时图像翻转而掩码没跟着翻转

现象:训练时验证集 Dice 始终上不去,可视化增强后的图像和掩码发现,掩码位置对不上图像内容。

原因:单独对 image 做翻转和旋转,label 只做了简单的类型转换,两个变换没有同步。模型学习的全是错位配对的特征。

解决:把 image 和 mask 放在同一个增强 pipeline 里,albumentations 的 Compose 会保证两者应用完全相同的变换参数。如果用 PyTorch 原生 transform,需要手动确保随机种子一致:

def random_flip_pair(image, mask): if np.random.rand() > 0.5: image = np.fliplr(image).copy() mask = np.fliplr(mask).copy() return image, mask

换成 .copy() 这一步也容易被忽略,np.fliplr 返回的是视图,后面转 tensor 时可能报错或出现诡异的内存错位。数据增强带来的问题往往在训练早期不暴露,等跑到第 50 个 epoch 才显现出来,排查成本极高,建议在增强管线完成后可视化三对样本再开训练。

6. 让结果可控的三个进阶习惯:测试时增强、patch 推理与输出合规

训练收敛之后,真正决定模型能不能落地的反而是推理阶段。这里分享三个我常用的技巧,代码量都不大但收益明显。

第一个是测试时增强。推理时对输入做水平翻转,把原图和翻转图的预测概率取平均后再 argmax。这个技巧几乎不增加训练成本,但能明显抑制模型对特定方向的偏好,尤其适合数据量少的医学场景:

def tta_predict(model, img, device): logits = torch.softmax(model(img.to(device)), dim=1) logits = logits + torch.softmax(model(torch.flip(img, dims=[3]).to(device)), dim=1) logits = logits / 2.0 return torch.argmax(logits, dim=1)

flip 的 dims=[3] 对应宽方向。个别数据集对翻转很敏感,TTA 不一定每次都涨点,跑一次对比就能判断值不值得保留。

第二个是大体积数据的 patch 推理。nii 体数据直接整图送入 UNet 会爆显存,常规做法是按滑窗裁剪成 256x256 patch 逐个推理。固定不重叠裁剪会在 patch 边界出现明显的拼接缝,解决方法是让 patch 之间有 50% 重叠,重叠区域的预测结果做平均,能有效消除边界伪影。这个技巧同时适用于 2D 切片和 3D 体数据。

第三个是推理结果的输出合规。医学分割的结果最终要给专业软件看,直接保存成掩码 png 会丢失空间位置信息。正确做法是找回训练前保存的 SimpleITK 图像对象,把预测数组塞回去再写 nii.gz,这样 spacing、origin、direction 全部保留:

pred_sitk = sitk.GetImageFromArray(pred_array.astype(np.uint8)) pred_sitk.CopyInformation(img_meta) # 复用原图的元信息 sitk.WriteImage(pred_sitk, "pred.nii.gz")

CopyInformation 会把 spacing 和方向一并拷贝,这点比手动 set_spacing 靠谱得多,不用逐项担心元信息丢失。

在医学图像分割上吃过最多次亏的地方不是网络结构,而是数据读取和标签映射。后来我养成了一个习惯:任何新数据集,都先把处理后的图像和掩码重叠可视化一遍,确认类别连续、方向正确、空间对齐,再开始调整模型。这种二十行代码的 sanity check,省下的时间远超训练本身的成本,希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询