
简介这份资源面向医学图像分割、语义分割与多类别分割的学习者和研究者提供一套基于U-Net的完整代码实现。U-Net凭借对称的收缩与扩展路径以及跳跃连接能在小样本数据下捕捉上下文信息并保留精细边界适合处理病灶定位、组织结构量化等任务。资源包共31个文件以8个py源码文件为核心涵盖模型定义、数据集加载、数据增强、训练与预测脚本并配有混淆矩阵计算模块另有14个pyc编译缓存、5个xml项目配置、1个txt依赖清单及readme说明整体约16KB结构紧凑便于快速上手。目前已有466人学习下载。读者可据此搭建训练与推理流程理解跳跃连接、多类别分割与评价指标实现并在此基础上结合注意力机制或残差结构做进一步优化。1. unet 医学图像分割从跑通到多类别落地的完整路径医学图像分割这个方向很多人第一次接触就是从 unet 开始的。原因很直接结构简单、论文清晰、代码复现门槛低肺部 CT、视网膜血管、细胞核、肝脏肿瘤这些公开数据集上都有现成的 baseline。但真正把它用到自己的数据上尤其是从二分类切到多类别分割时翻车点会集中爆发——标签对不上、类别不均衡、Dice 卡在 0.6 上不去、预测出来全是背景。这篇笔记就围绕 unet 医学图像分割、语义分割、多类别分割代码这条主线把数据准备、模型搭建、训练调参、多类别改造、推理验证整条链路拆开讲。适合已经能跑通 demo、准备上自己数据集的人也适合想搞清楚多类别语义分割和实例分割区别的从业者。读完你应该能独立搭出一套可复现的多类别 unet 训练流程并且知道每一步参数为什么这么设。2. 数据准备与标签体系多类别分割的地基2.1 医学图像为什么不能直接套自然图像预处理医学图像和自然图像最大的差别在于灰度分布和对比度。CT 的 HU 值范围大概在 -1000 到 3000直接归一化到 [0,1] 会把软组织信息压扁MRI 又没有固定物理量纲不同扫描仪、不同序列的强度差异很大。我一般会先做窗宽窗位截断再做 z-score 归一化而不是无脑除以 255。import numpy as np def ct_window_normalize(img, window_center40, window_width400): # 按窗宽窗位截断保留软组织对比度 low window_center - window_width // 2 high window_center window_width // 2 img np.clip(img, low, high) # 再归一化到 [0,1] img (img - low) / (high - low) return img.astype(np.float32)这段代码的关键在 window_center 和 window_width 两个参数。腹部软组织常用 40/400肺部常用 -600/1500骨窗用 300/1500。选错窗口模型看到的就是一片灰Dice 自然上不去。如果你的数据是 MRI跳过窗宽窗位直接做百分位截断比如 1% 到 99%再 z-score 更稳。2.2 多类别标签的三种常见格式与转换多类别分割的标签格式直接决定损失函数怎么写。常见有三种一是每类一个二值 maskone-hot 堆叠二是单通道灰度图像素值就是类别 id0 背景、1 类 A、2 类 B三是 RGB 彩色标注图。unet 多类别训练最省事的是第二种因为 CrossEntropyLoss 和 DiceLoss 都能直接吃。格式存储优点缺点one-hot 多通道N 个 PNG直观占空间读取慢单通道 id 图1 个 PNG省空间直接喂损失需要保证 id 连续RGB 彩色图1 个 PNG标注工具友好必须做颜色到 id 的映射从 RGB 转 id 图的代码很常见但坑在颜色映射表必须和标注规范严格一致差一个像素值就会多出一个类别。import numpy as np # 颜色到类别 id 的映射必须和标注规范一致 COLOR_MAP { (0, 0, 0): 0, # 背景 (255, 0, 0): 1, # 类 A (0, 255, 0): 2, # 类 B (0, 0, 255): 3, # 类 C } def rgb_to_id(mask_rgb): id_map np.zeros(mask_rgb.shape[:2], dtypenp.uint8) for color, cid in COLOR_MAP.items(): match np.all(mask_rgb color, axis-1) id_map[match] cid return id_map转换完一定要做一次校验统计 id_map 里出现的唯一值看是否和预期类别数一致。我见过标注工具导出时做了抗锯齿边缘出现 (128,0,0) 这种中间色结果整张图多出几十个伪类别训练直接崩。2.3 类别不均衡多类别分割绕不开的第一道坎医学数据里背景通常占 90% 以上小病灶可能只占 0.5%。如果直接上 CrossEntropyLoss模型学会全预测背景就能拿到很高的 accuracy但 Dice 接近 0。常见做法有三种加权 CrossEntropy、Dice Loss、以及两者的组合。我一般用 0.5 倍 CrossEntropy 加 0.5 倍多类别 Dice权重按类别频率的倒数开方来设比直接取倒数温和不容易过拟合小类。import torch import torch.nn as nn class MultiClassDiceLoss(nn.Module): def __init__(self, num_classes, smooth1e-5): super().__init__() self.num_classes num_classes self.smooth smooth def forward(self, logits, targets): # logits: [B, C, H, W], targets: [B, H, W] probs torch.softmax(logits, dim1) targets_onehot torch.nn.functional.one_hot( targets, self.num_classes).permute(0, 3, 1, 2).float() dims (0, 2, 3) inter (probs * targets_onehot).sum(dims) union probs.sum(dims) targets_onehot.sum(dims) dice (2 * inter self.smooth) / (union self.smooth) return 1 - dice.mean()smooth 参数别设太小1e-5 是经验值太小在空类别上会出 NaN。num_classes 必须包含背景这是新手最常犯的错——把背景漏掉通道数对不上报错还看不懂。3. unet 网络搭建与多类别输出改造3.1 经典 unet 结构里哪些部分值得改原始 unet 是 2015 年为二分类细胞分割设计的四次下采样、四次上采样、通道数 64 起步翻倍到 1024。放到多类别医学分割上有三处我一般会动一是把输出通道从 1 改成 num_classes二是把最后的 sigmoid 换成 softmax或者干脆不激活交给损失函数三是把 BatchNorm 换成 InstanceNorm 或 GroupNorm因为医学数据 batch size 往往只能开到 2 到 4BN 统计量不稳。import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch, normgroup): super().__init__() if norm group: n nn.GroupNorm(8, out_ch) else: n nn.BatchNorm2d(out_ch) self.block nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1, biasFalse), n, nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1, biasFalse), n, nn.ReLU(inplaceTrue), ) def forward(self, x): return self.block(x)GroupNorm 的组数一般取 8 或 16要求能整除通道数。如果通道数是 64取 8 没问题如果自定义了 48 通道取 8 也能整除。这个细节不注意会直接抛异常。3.2 多类别输出的通道与激活怎么配多类别语义分割的输出层通道数等于类别数含背景激活函数的选择取决于损失函数。用 CrossEntropyLoss 时网络输出 logits不要加 softmax因为 PyTorch 的 CE 内部会做 log_softmax用 Dice Loss 时需要先 softmax 再算。两者组合时把 logits 同时喂给 CE 和 DiceDice 内部自己做 softmax这样最干净。class UNet(nn.Module): def __init__(self, in_ch1, num_classes4, base64): 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) self.up4 nn.ConvTranspose2d(base * 16, base * 8, 2, stride2) self.dec4 DoubleConv(base * 16, base * 8) self.up3 nn.ConvTranspose2d(base * 8, base * 4, 2, stride2) self.dec3 DoubleConv(base * 8, base * 4) self.up2 nn.ConvTranspose2d(base * 4, base * 2, 2, stride2) self.dec2 DoubleConv(base * 4, base * 2) self.up1 nn.ConvTranspose2d(base * 2, base, 2, stride2) self.dec1 DoubleConv(base * 2, base) 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], 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(d1) # 返回 logits不加 softmaxbase 通道数从 64 起步是经典配置显存不够可以降到 32 或 16但别低于 16否则特征表达能力不够。num_classes 一定要数清楚背景算一类这是多类别分割代码里最容易数错的地方。3.3 训练循环与关键参数设置训练循环本身不复杂关键是几个参数学习率、优化器、batch size、以及验证指标。我一般用 AdamW学习率 1e-3 配 cosine 退火batch size 能开多大开多大开不大就用梯度累积。验证指标用每类 Dice 加平均 Dice光看 loss 会被背景主导。from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR model UNet(in_ch1, num_classes4).cuda() optimizer AdamW(model.parameters(), lr1e-3, weight_decay1e-4) scheduler CosineAnnealingLR(optimizer, T_max100) ce_loss nn.CrossEntropyLoss(weightclass_weights.cuda()) dice_loss MultiClassDiceLoss(num_classes4) for epoch in range(100): model.train() for img, mask in train_loader: img, mask img.cuda(), mask.cuda() logits model(img) loss 0.5 * ce_loss(logits, mask) 0.5 * dice_loss(logits, mask) optimizer.zero_grad() loss.backward() optimizer.step() scheduler.step()class_weights 按类别频率倒数开方算别用原始倒数容易过拟合小类。T_max 设成总 epoch 数cosine 退火到接近 0比 step 调度稳。梯度累积的话把 loss 除以累积步数再 backward别直接累加。4. 多类别分割的避坑与排查清单4.1 现象训练 loss 正常下降但 Dice 一直是 0原因通常是标签 id 和输出通道对不上。比如标签里背景是 255 而不是 0或者类别 id 从 1 开始但 num_classes 设成了类别数没加背景。解决方法是训练前打印标签的唯一值和最大值确认 id 范围是 [0, num_classes-1]。我一般会在 Dataset 里加一句断言id 超范围直接报错比训练到一半才发现强。4.2 现象某些类别 Dice 始终为 0其他类别正常这是典型的小类被背景淹没。原因可能是该类样本太少或者损失权重没设对。解决方法是先统计每类像素占比占比低于 0.1% 的类别考虑过采样含该类的 patch或者把损失权重调高。另外检查一下该类在验证集里是否真的存在有时候是数据划分问题验证集里根本没这个类。4.3 现象预测结果边缘锯齿严重小结构丢失原因通常是下采样太深或者上采样用了最近邻。unet 四次下采样对 512x512 的图还行对 256x256 的图就有点深了小结构在 bottleneck 处信息丢光。解决方法是减少下采样次数或者把 ConvTranspose2d 换成双线性插值加卷积边缘更平滑。另外可以在 skip connection 上加注意力门控让网络聚焦小结构。4.4 现象验证集 Dice 比训练集低很多过拟合的典型表现。医学数据量小几百张图很常见。解决方法是加数据增强弹性形变、随机旋转、亮度扰动加 dropout 或 weight decay以及早停。弹性形变对医学图像特别有效因为器官形状本身就有自然变化。我一般用 albumentations 的 ElasticTransformalpha 设 1sigma 设 50别太猛。4.5 现象多卡训练时 Dice 计算出现 NaN原因通常是某张卡上某个类别在当前 batch 里一个像素都没有Dice 分母为 0。解决方法是 Dice Loss 里加 smooth 项或者在计算前判断 union 是否为 0是就跳过该类。分布式训练时还要注意 all_reduce 的同步别各卡各算。5. 推理、验证与多类别分割的进阶技巧推理阶段最容易忽略的是滑窗和重叠。医学图像往往比显存能容纳的尺寸大直接 resize 会丢细节。我一般用 512x512 的滑窗步长 256重叠区域取概率平均比直接取最大类别更稳。多类别输出先 softmax 得到每类概率再在重叠区累加最后 argmax。def sliding_window_inference(model, image, window512, stride256, num_classes4): model.eval() _, h, w image.shape prob_map torch.zeros((num_classes, h, w), deviceimage.device) count_map torch.zeros((1, h, w), deviceimage.device) for y in range(0, h, stride): for x in range(0, w, stride): y1, x1 min(y, h - window), min(x, w - window) patch image[:, y1:y1window, x1:x1window].unsqueeze(0) with torch.no_grad(): logits model(patch) probs torch.softmax(logits, dim1).squeeze(0) prob_map[:, y1:y1window, x1:x1window] probs count_map[:, y1:y1window, x1:x1window] 1 prob_map / count_map.clamp(min1) return prob_map.argmax(dim0)window 和 stride 的比值决定重叠程度一般 stride 取 window 的一半。太小推理慢太大边缘会有拼接痕迹。num_classes 要和训练时严格一致顺序也不能变否则类别会错位。验证多类别分割不能只看平均 Dice要每类单独看。我习惯画一个混淆矩阵看哪些类之间容易混。医学图像里相邻器官边界模糊混淆很正常这时候可以考虑在损失里加边界加权或者后处理用 CRF refine。另外如果类别数超过 5 个建议把背景单独拿出来看背景 Dice 通常虚高会拉高平均值掩盖问题。最后一个习惯每次改完数据或模型先跑 5 个 epoch 的小实验看 loss 和 Dice 的趋势趋势不对就别浪费算力跑满。这个习惯帮我省了无数次通宵。希望帮到你。本文还有配套的精品资源点击获取