
简介这份资源面向医学图像分割、语义分割与多类别分割的学习者与研究者提供一套可直接运行的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 ReLUunet 的基本单元 def __init__(self, in_ch, out_ch): super().__init__() self.net nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.net(x) 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) # 解码器转置卷积上采样后 concat 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) # 输出层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], 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训练时交给损失函数这段代码里几个参数需要重点说明。in_ch对应输入模态数单模态 CT 是 1RGB 病理图是 3多模态 MRI 比如 T1T2 可以设成 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(weighttorch.tensor([0.1, 1.0, 1.0, 1.0])) # 假设模型输出 [B, 4, H, W]标签 [B, H, W] 且为 long model UNet(in_ch1, num_classesnum_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. 训练自己的医学图像数据集从标注格式到 dataloader3.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_size256, augmentFalse): 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], interpolationTF.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_size256, augmentTrue) dl DataLoader(ds, batch_size8, shuffleTrue, num_workers4) 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 输入、base64 的 unet8GB 显存大概能跑 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(dim1) # [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_ch1, num_classes4, base64).to(device) optimizer torch.optim.AdamW(model.parameters(), lr1e-3, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max100) criterion nn.CrossEntropyLoss(weighttorch.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(fepoch {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设成 N1要么把标签减 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(modebilinear)加一个 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_size256, overlap32, num_classes4): 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:y1patch_size, x1:x1patch_size].unsqueeze(0) logits model(patch) prob torch.softmax(logits, dim1).squeeze(0) prob_map[:, y1:y1patch_size, x1:x1patch_size] prob count_map[:, y1:y1patch_size, x1:x1patch_size] 1 prob_map / count_map.clamp(min1) return prob_map.argmax(dim0)overlap设 32 到 64 之间太小拼接处会有明显接缝太大推理慢。patch_size要和训练时一致训练用 256 推理也用 256否则 BN 统计量不匹配。拼接时用概率图累加再平均比直接取类别图拼接更平滑边缘过渡自然。5.2 后处理连通域过滤与形态学操作模型输出的类别图往往有零星小噪点或者某个器官被分成几块。后处理能明显改善视觉效果。import numpy as np from scipy import ndimage def postprocess(pred_mask, num_classes4, min_size100): 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, structurenp.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再给混淆矩阵才能看出模型到底哪里弱。类别DiceIoUPrecisionRecall背景0.980.960.990.97肝脏0.920.850.900.94肾脏0.870.770.890.85脾脏0.810.680.840.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 值确认没有形状不匹配和标签越界再开完整训练。这个检查花不了一分钟但能省下几小时白跑的训练。医学图像分割这个方向模型结构其实不是瓶颈数据质量和训练细节才是拉开差距的地方。希望帮到你。本文还有配套的精品资源点击获取