简介:这份资源面向医学图像分割、语义分割与多类别分割的学习者与研究者,提供一套可直接运行的U-Net实现代码,帮助解决小数据集下分割精度不足、边界细节丢失等常见问题。压缩包共31个文件,约16KB,以8个Python源码文件为核心,涵盖模型定义、数据集加载、数据增强、训练与预测脚本,并配有混淆矩阵评估模块;另有14个pyc缓存文件、5个xml与iml等IDE配置、readme及requirements说明,目录结构清晰,便于快速复现与二次开发。资源围绕U-Net的收缩路径、扩展路径与跳跃连接展开,可迁移至医学病灶定位、组织结构量化及通用语义分割任务。目前已有466人学习下载,适合希望掌握分割网络训练流程、评估指标与工程组织方式的中级读者参考。
1. unet 医学图像分割代码:从跑通到多类别落地的真实路径
医学图像分割这个方向,很多人第一次接触就是拿一份 unet 代码在公开数据集上跑一遍,看着 dice 从 0.3 涨到 0.85,觉得不过如此。真正上手自己科室或自己项目的数据时才发现,单类别二分类的 unet 和能处理肝脏、肾脏、脾脏、胰腺多类别分割的 unet,中间隔着一整套数据管线、损失函数选型和后处理逻辑。这份笔记就围绕 unet 医学图像分割、语义分割、多类别分割代码这条主线,把从环境搭建、数据组织、模型改造、训练调参到推理部署的完整链路拆开讲。适合已经看过 unet 结构图、想真正把语义分割模型跑在自己数据上的工程师和研究生,也适合做过二分类分割、想扩展到多类别场景的从业者。下面所有代码都是可复现的最小实现,参数含义和踩坑点会逐个说明。
2. unet 语义分割代码骨架:编码器、解码器与跳跃连接怎么落地
2.1 为什么医学图像分割偏爱 unet 这套结构
语义分割算法里,unet 之所以在医学图像领域站住脚,核心在于它的跳跃连接把浅层的高分辨率细节和深层的语义信息拼在一起。医学图像比如 CT、MRI、超声,边界往往模糊,器官之间灰度接近,纯靠深层特征上采样回来,边缘会糊成一片。unet 的编码器逐层下采样提取语义,解码器逐层上采样恢复分辨率,每次上采样后把编码器对应层的特征图 concat 进来,让解码器在恢复空间细节时有据可依。
从代码角度看,一个标准 unet 由三部分组成:下采样块(DoubleConv + MaxPool)、瓶颈层、上采样块(Upsample 或转置卷积 + concat + DoubleConv)。医学图像通常尺寸不大,比如 512×512 的 CT 切片,下采样四次到 32×32 的瓶颈层已经足够。如果输入是 3D 体数据,常见做法是把 2D unet 的卷积核换成 3D 卷积,或者沿 z 轴切片当 2D 处理再拼接,前者显存吃紧,后者会丢失层间连续性,选哪种取决于你的标注是不是逐层做的。
多类别分割和单类别的区别不在网络结构本身,而在输出通道数和损失函数。单类别输出 1 个通道配 sigmoid,多类别输出 N 个通道配 softmax,背景也算一类。很多人第一次改多类别时忘了把背景算进去,导致类别数少一,训练时 loss 一直不降,这是血泪经验里最常见的一条。
2.2 一份可直接运行的最小 unet 代码
下面这份代码是 2D unet 的最小实现,支持任意类别数,输入通道可配置,适合作为医学图像分割代码的起点。
import torch import torch.nn as nn import torch.nn.functional as F class DoubleConv(nn.Module): """两次 3x3 卷积 + BN + ReLU,unet 的基本单元""" def __init__(self, in_ch, out_ch): super().__init__() self.net = nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) def forward(self, x): return self.net(x) class UNet(nn.Module): def __init__(self, in_ch=1, num_classes=4, base=64): super().__init__() # 编码器:每次下采样通道翻倍 self.enc1 = DoubleConv(in_ch, base) self.enc2 = DoubleConv(base, base*2) self.enc3 = DoubleConv(base*2, base*4) self.enc4 = DoubleConv(base*4, base*8) self.pool = nn.MaxPool2d(2) # 瓶颈层 self.bottleneck = DoubleConv(base*8, base*16) # 解码器:转置卷积上采样后 concat self.up4 = nn.ConvTranspose2d(base*16, base*8, 2, stride=2) self.dec4 = DoubleConv(base*16, base*8) self.up3 = nn.ConvTranspose2d(base*8, base*4, 2, stride=2) self.dec3 = DoubleConv(base*8, base*4) self.up2 = nn.ConvTranspose2d(base*4, base*2, 2, stride=2) self.dec2 = DoubleConv(base*4, base*2) self.up1 = nn.ConvTranspose2d(base*2, base, 2, stride=2) self.dec1 = DoubleConv(base*2, base) # 输出层:num_classes 个通道,不接 softmax self.out = nn.Conv2d(base, 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], dim=1)) d3 = self.dec3(torch.cat([self.up3(d4), e3], dim=1)) d2 = self.dec2(torch.cat([self.up2(d3), e2], dim=1)) d1 = self.dec1(torch.cat([self.up1(d2), e1], dim=1)) return self.out(d1) # 返回 logits,训练时交给损失函数这段代码里几个参数需要重点说明。in_ch对应输入模态数,单模态 CT 是 1,RGB 病理图是 3,多模态 MRI 比如 T1+T2 可以设成 2。num_classes是包含背景的总类别数,做肝脏、肾脏、脾脏三器官分割时,背景+3 个器官等于 4。base是基础通道数,64 是常规起点,显存不够可以降到 32,但要注意降太多会影响小目标分割精度。
输出层故意不接 softmax,因为 PyTorch 的CrossEntropyLoss内部已经包含 log_softmax,如果模型里再加一层 softmax,等于做了两次归一化,梯度会异常,训练 loss 会卡住不动。这是 unet 使用时的注意事项里排前三的坑。
2.3 多类别分割的输出通道与损失函数怎么配
多类别语义分割的标签组织方式和单类别完全不同。单类别标签是 0/1 二值图,多类别标签是每个像素的类别索引,背景为 0,器官 1 到 N-1。标签必须是long类型,不能是float,否则CrossEntropyLoss会直接报错。
import torch.nn as nn # 多类别分割标准配置 num_classes = 4 # 背景 + 3 个器官 criterion = nn.CrossEntropyLoss(weight=torch.tensor([0.1, 1.0, 1.0, 1.0])) # 假设模型输出 [B, 4, H, W],标签 [B, H, W] 且为 long model = UNet(in_ch=1, num_classes=num_classes) logits = model(torch.randn(2, 1, 256, 256)) target = torch.randint(0, num_classes, (2, 256, 256)).long() loss = criterion(logits, target) print(loss.item())weight参数是给每个类别加权的,医学图像里背景像素通常占 90% 以上,不加权的话模型会倾向于全预测背景,dice 看起来还行但器官一个都分不出来。常见做法是背景权重设 0.1 到 0.3,前景类别设 1.0,具体值根据你的类别像素占比调。如果某个器官特别小,比如胰腺,权重可以再往上提到 2.0 甚至 3.0。
另一个常见组合是CrossEntropyLoss加DiceLoss,两者按 0.5:0.5 加权。CE 负责像素级分类稳定收敛,Dice 负责优化区域重叠度,对类别不平衡更鲁棒。DiceLoss 在多类别下要对每个类别单独算 dice 再平均,不能把所有类别混在一起算。
3. 训练自己的医学图像数据集:从标注格式到 dataloader
3.1 医学图像分割数据集制作的三种常见格式
拿到一批 CT 或 MRI 数据后,第一步是搞清楚标注是什么格式。常见的有三种:PNG 掩码图、NIfTI 文件、COCO JSON。PNG 掩码最简单,每个像素值就是类别索引,适合 2D 切片。NIfTI 是 3D 体数据格式,.nii或.nii.gz,标注和图像在同一个空间坐标系下,适合 3D 分割。COCO JSON 多用于自然图像,医学图像里少见,但有些标注工具会导出这个格式。
如果标注是 NIfTI,用nibabel读取,注意方向矩阵和 spacing。不同设备导出的 NIfTI 方向可能不一致,直接切片会导致左右颠倒。常见做法是用nib.as_closest_canonical()统一到 RAS 方向再处理。如果标注是 PNG,要确认像素值是不是从 0 开始连续,有些工具导出时背景是 255,器官是 1、2、3,这种要先做映射。
数据集划分上,医学图像不能随机按切片划分,因为同一患者的相邻切片高度相似,随机划分会导致训练集和验证集泄漏,验证 dice 虚高。正确做法是按患者划分,同一患者的所有切片只出现在一个集合里。这一点在公开数据集上不明显,但在自己数据上如果不注意,模型上线后性能会断崖式下跌。
3.2 自定义 Dataset 与多类别标签处理
下面是一个支持多类别分割的 Dataset 实现,输入是图像文件夹和掩码文件夹,按文件名配对。
import os import numpy as np import torch from torch.utils.data import Dataset, DataLoader from PIL import Image import torchvision.transforms.functional as TF class MedSegDataset(Dataset): def __init__(self, img_dir, mask_dir, img_size=256, augment=False): self.img_dir = img_dir self.mask_dir = mask_dir self.img_size = img_size self.augment = augment # 只保留有对应掩码的样本 self.names = [f for f in os.listdir(img_dir) if os.path.exists(os.path.join(mask_dir, f))] def __len__(self): return len(self.names) def __getitem__(self, idx): name = self.names[idx] img = Image.open(os.path.join(self.img_dir, name)).convert('L') mask = Image.open(os.path.join(self.mask_dir, name)) # 统一尺寸,分割任务必须用最近邻插值处理掩码 img = TF.resize(img, [self.img_size, self.img_size]) mask = TF.resize(mask, [self.img_size, self.img_size], interpolation=TF.InterpolationMode.NEAREST) img = TF.to_tensor(img) # [1, H, W],归一化到 0-1 mask = torch.from_numpy(np.array(mask)).long() # [H, W] # 数据增强:图像和掩码必须同步变换 if self.augment and torch.rand(1) > 0.5: img = TF.hflip(img) mask = TF.hflip(mask.unsqueeze(0)).squeeze(0) return img, mask # 使用示例 ds = MedSegDataset('./data/images', './data/masks', img_size=256, augment=True) dl = DataLoader(ds, batch_size=8, shuffle=True, num_workers=4) imgs, masks = next(iter(dl)) print(imgs.shape, masks.shape, masks.dtype) # [8,1,256,256] [8,256,256] torch.int64这段代码有几个关键点。掩码 resize 必须用NEAREST插值,用双线性会把类别索引插成小数,long()转换后类别全乱。数据增强时图像和掩码必须用同一组随机参数,上面用同一个随机数控制翻转,实际项目里建议用albumentations或monai的同步增强接口,避免手写时漏掉某个变换。
num_workers在 Windows 上设大于 0 可能报错,Linux 上一般设 4 到 8。如果数据集不大,设 0 也能跑,只是慢。batch_size受显存限制,256×256 输入、base=64 的 unet,8GB 显存大概能跑 batch 8 到 12。
3.3 训练循环与验证指标 dice 的计算
训练循环本身不复杂,但多类别分割的验证指标计算容易写错。dice 要按类别分别算再平均,不能把所有前景混在一起。
import torch import torch.nn as nn from tqdm import tqdm def dice_per_class(pred, target, num_classes): """pred: [B,C,H,W] logits, target: [B,H,W] long""" pred = pred.argmax(dim=1) # [B,H,W] dice_list = [] for c in range(1, num_classes): # 跳过背景 p = (pred == c) t = (target == c) inter = (p & t).sum().float() union = p.sum().float() + t.sum().float() if union == 0: continue # 该类别在本 batch 不存在 dice_list.append(2 * inter / union) return torch.stack(dice_list).mean() if dice_list else torch.tensor(0.0) def train_one_epoch(model, loader, optimizer, criterion, device, num_classes): model.train() total_loss = 0 for img, mask in tqdm(loader): img, mask = img.to(device), mask.to(device) optimizer.zero_grad() logits = model(img) loss = criterion(logits, mask) loss.backward() optimizer.step() total_loss += loss.item() return total_loss / len(loader) # 主训练配置 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = UNet(in_ch=1, num_classes=4, base=64).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=100) criterion = nn.CrossEntropyLoss(weight=torch.tensor([0.2, 1.0, 1.0, 1.0]).to(device)) for epoch in range(100): loss = train_one_epoch(model, dl, optimizer, criterion, device, 4) scheduler.step() if epoch % 10 == 0: print(f'epoch {epoch}, loss {loss:.4f}')AdamW比Adam多了正确的权重衰减实现,医学图像数据量小,正则化很重要。学习率 1e-3 是起点,如果 loss 震荡明显降到 3e-4。CosineAnnealingLR让学习率按余弦曲线下降,比阶梯下降更平滑,适合分割任务。验证时记得model.eval()加torch.no_grad(),否则 BN 层统计量会被验证数据污染。
4. unet 多类别分割训练避坑:从 loss 不降到显存爆炸
4.1 标签值越界导致 loss 直接 NaN
现象:训练第一个 batch 就报CUDA error: device-side assert triggered或者 loss 变成 NaN。
原因:CrossEntropyLoss要求 target 的值在[0, num_classes-1]范围内。如果掩码图里背景是 255,或者标注工具导出时用了 1、2、3、4 而num_classes设成了 4,索引 4 就越界了。
解决:训练前先统计掩码的唯一值,np.unique(mask)打印出来确认。如果背景是 255,做一次映射mask[mask == 255] = 0。如果类别是 1 到 N,要么把num_classes设成 N+1,要么把标签减 1 映射到 0 到 N-1。这个检查建议写进 Dataset 的__init__里,跑一次全量统计,别等训练时报错。
4.2 背景权重设太高导致器官全丢
现象:训练 loss 降得很快,验证 dice 也有 0.7 左右,但可视化一看全是背景,器官一个都没分出来。
原因:背景像素占比太高,如果背景权重没降下来,模型发现全预测背景就能拿到很低的 loss,直接躺平。dice 指标因为背景占大头,看起来也不低。
解决:给CrossEntropyLoss的weight参数里背景设小值,比如 0.1 到 0.3。更稳妥的做法是加 DiceLoss 联合训练,Dice 对类别不平衡不敏感。另外验证时一定要按类别打印 dice,只看平均 dice 会被背景拉高。如果某个类别 dice 一直是 0,说明模型完全没学到,回去检查标签和权重。
4.3 显存爆炸与 batch size 的取舍
现象:训练到一半报CUDA out of memory,或者一开始就爆。
原因:unet 的显存占用和输入尺寸、base 通道数、batch size 都成正比。512×512 输入比 256×256 显存占用大约 4 倍。base 从 64 提到 128,参数量和显存都翻倍。
解决:优先降 batch size 到 2 甚至 1,配合梯度累积模拟大 batch。如果还不够,把 base 降到 32,或者用混合精度训练。混合精度在 PyTorch 里用torch.cuda.amp几行就能加上,显存能省 30% 到 40%,速度也快。注意混合精度下 loss 要放在GradScaler里 scale,否则梯度会下溢。
scaler = torch.cuda.amp.GradScaler() for img, mask in loader: with torch.cuda.amp.autocast(): logits = model(img) loss = criterion(logits, mask) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() optimizer.zero_grad()4.4 验证集 dice 虚高但推理效果差
现象:训练时验证 dice 0.9,拿模型去推理新数据,分割结果一塌糊涂。
原因:最常见的是数据泄漏,同一患者的切片同时出现在训练集和验证集。其次是验证时用了训练集的归一化参数,而推理时新数据的灰度分布不同。医学图像不同设备、不同扫描协议的灰度范围差异很大,训练时如果按数据集全局均值方差归一化,推理时必须用同一组参数。
解决:按患者划分数据集,写个脚本按患者 ID 分组再分。归一化参数保存下来,推理时加载同一组。如果新数据分布差异大,考虑用直方图匹配或者自适应归一化。另外验证时加model.eval(),别让 dropout 和 BN 在验证时还起作用。
4.5 上采样方式选转置卷积还是双线性插值
现象:转置卷积训练时出现棋盘格伪影,分割边缘有规律性网格。
原因:转置卷积的卷积核步长大于 1 时,如果核大小不能被步长整除,输出会有重叠不均匀,产生棋盘效应。
解决:把转置卷积换成nn.Upsample(mode='bilinear')加一个 1×1 卷积调整通道,或者用nn.ConvTranspose2d时确保kernel_size能被stride整除,比如 kernel 4 stride 2。医学图像分割对边缘敏感,双线性插值加卷积的组合更稳,虽然参数量少一点,但伪影问题基本没有。
5. 推理部署与多类别分割后处理:让模型真正能用
5.1 滑窗推理处理大尺寸医学图像
医学图像原始尺寸往往超过模型输入,比如全切片病理图可能上万像素,CT 也是 512×512 但需要处理整个 3D 体积。常见做法是滑窗推理,把大图切成有重叠的 patch,逐块预测后拼接。
import torch import numpy as np @torch.no_grad() def sliding_window_inference(model, image, patch_size=256, overlap=32, num_classes=4): """image: [1, H, W] tensor, 返回 [H, W] 类别图""" model.eval() _, H, W = image.shape stride = patch_size - overlap prob_map = torch.zeros(num_classes, H, W) count_map = torch.zeros(1, H, W) for y in range(0, H, stride): for x in range(0, W, stride): y1, x1 = min(y, H - patch_size), min(x, W - patch_size) y1, x1 = max(y1, 0), max(x1, 0) patch = image[:, y1:y1+patch_size, x1:x1+patch_size].unsqueeze(0) logits = model(patch) prob = torch.softmax(logits, dim=1).squeeze(0) prob_map[:, y1:y1+patch_size, x1:x1+patch_size] += prob count_map[:, y1:y1+patch_size, x1:x1+patch_size] += 1 prob_map /= count_map.clamp(min=1) return prob_map.argmax(dim=0)overlap设 32 到 64 之间,太小拼接处会有明显接缝,太大推理慢。patch_size要和训练时一致,训练用 256 推理也用 256,否则 BN 统计量不匹配。拼接时用概率图累加再平均,比直接取类别图拼接更平滑,边缘过渡自然。
5.2 后处理:连通域过滤与形态学操作
模型输出的类别图往往有零星小噪点,或者某个器官被分成几块。后处理能明显改善视觉效果。
import numpy as np from scipy import ndimage def postprocess(pred_mask, num_classes=4, min_size=100): """pred_mask: [H, W] numpy 类别图""" result = pred_mask.copy() for c in range(1, num_classes): binary = (pred_mask == c) labeled, n = ndimage.label(binary) for i in range(1, n + 1): if (labeled == i).sum() < min_size: result[labeled == i] = 0 # 小于阈值的连通域归背景 # 闭运算填补小孔 for c in range(1, num_classes): mask_c = (result == c) mask_c = ndimage.binary_closing(mask_c, structure=np.ones((3, 3))) result[mask_c & (result == 0)] = c return resultmin_size根据你的目标大小定,比如肝脏在 256×256 图上大概几千像素,设 100 到 500 都合理。闭运算的structure用 3×3 全 1 就行,太大容易把相邻器官粘在一起。后处理不是必须的,但如果你的分割结果要给人看或者做体积测量,加上会专业很多。
5.3 多类别分割的评估报告怎么写
评估不能只看一个平均 dice。按类别列出 dice、IoU、precision、recall,再给混淆矩阵,才能看出模型到底哪里弱。
| 类别 | Dice | IoU | Precision | Recall |
|---|---|---|---|---|
| 背景 | 0.98 | 0.96 | 0.99 | 0.97 |
| 肝脏 | 0.92 | 0.85 | 0.90 | 0.94 |
| 肾脏 | 0.87 | 0.77 | 0.89 | 0.85 |
| 脾脏 | 0.81 | 0.68 | 0.84 | 0.78 |
从这张表能看出脾脏最弱,可能因为脾脏边界模糊或者训练样本少。下一步要么补脾脏标注,要么给脾脏类别加权。混淆矩阵能进一步看出脾脏被误分成什么,如果大量脾脏像素被分成肝脏,说明两个器官特征太接近,考虑加多模态输入或者换更强的编码器。
6. 把 unet 从能跑推到好用:三个我反复验证过的技巧
第一个技巧是深监督。在解码器每一层上采样后接一个 1×1 卷积输出辅助预测,和最终输出一起算 loss,辅助 loss 权重设 0.3 到 0.4。这个改动几乎不增加推理成本,但训练时梯度能更直接地传到浅层,收敛更快,小目标分割 dice 通常能涨 2 到 3 个点。代码上就是在forward里把 d4、d3、d2 各接一个nn.Conv2d(base*8, num_classes, 1)之类的头,训练时把这些辅助输出和主输出一起算 CE loss 求和。推理时只用主输出,辅助头丢掉。
第二个技巧是学习率 warmup 加余弦退火。医学图像数据集小,一开始用 1e-3 学习率容易震荡,前 5 个 epoch 从 1e-5 线性升到 1e-3,再余弦降到 1e-6。这个组合我在多个分割任务上试过,比固定学习率稳定得多,最终 dice 也高一点。实现上用torch.optim.lr_scheduler.LambdaLR写个 warmup 函数,再接CosineAnnealingLR。
第三个技巧是测试时增强。推理时把输入做水平翻转、垂直翻转、旋转 90 度,各预测一次,概率图平均后再取 argmax。这个操作推理时间翻 4 倍,但 dice 通常能涨 1 到 2 个点,对边界模糊的器官尤其明显。如果推理延迟不敏感,比如离线分析场景,值得加上。代码上就是把sliding_window_inference包一层,对每个变换后的输入推理再逆变换回来累加。
最后说个我自己的习惯:每次改完模型或数据管线,先拿一个 batch 过一遍,打印输入输出形状、标签唯一值、loss 值,确认没有形状不匹配和标签越界,再开完整训练。这个检查花不了一分钟,但能省下几小时白跑的训练。医学图像分割这个方向,模型结构其实不是瓶颈,数据质量和训练细节才是拉开差距的地方。希望帮到你。
本文还有配套的精品资源,点击获取