ARTICLE DETAIL

资讯详情

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

U-Net医学图像分割代码包实战:从跑通到多类别调优

U-Net医学图像分割代码包实战:从跑通到多类别调优 简介这份资源面向医学图像分割、语义分割与多类别分割的学习者和研究者提供一套基于U-Net的完整代码实现。U-Net凭借对称的收缩与扩展路径以及跳跃连接能在小样本数据下捕捉上下文信息并保留精细边界适合疾病诊断、病变定位与组织结构量化等场景。压缩包共31个文件约16KB以8个py源码文件为核心涵盖模型定义、数据集加载、数据增强、训练与预测脚本并配有混淆矩阵评估模块另有14个pyc缓存、5个xml与iml等IDE配置、readme和requirements说明目录结构清晰便于直接运行与二次修改。目前已有466人学习。读者可据此快速搭建训练与推理流程理解跳跃连接、多类别分割与评估指标的具体实现并在此基础上结合注意力机制或残差结构做进一步优化。1. 拿到一份 U-Net 分割代码包先别急着 train.py很多人第一次接触医学图像分割是从一份 U-Net 代码包开始的。解压之后看到train.py、predict.py、model.py、dataset.py、transforms.py、confuse_matrix.py这一串文件第一反应往往是直接python train.py然后被路径报错、通道数不匹配、显存溢出轮番教育。这份代码包的价值不在于它实现了 U-Net 这个 2015 年就提出的对称编解码结构而在于它把医学图像分割、语义分割、多类别分割三条任务线用同一套骨架串了起来收缩路径负责抓上下文扩展路径配合跳跃连接恢复分辨率dataset.py和transforms.py负责把原始图像和标签喂成网络能吃的张量confuse_matrix.py负责在训练后告诉你每个类别到底分对了多少。它适合谁手里有带标注的医学影像或语义分割数据集、想跑通一个能改能调的多类别分割基线的人。不适合想直接拿预训练权重出结果的人因为包里没有权重文件训练得自己来。下面按「这份代码怎么跑起来 → 每个模块在干什么 → 多类别怎么配 → 坑在哪 → 怎么验证」的顺序拆开讲每一步都落到能抄的参数和命令上。2. 把代码包跑起来环境、目录与第一次前向2.1 依赖安装与目录约定requirements.txt是这份代码包的入口清单常见内容是 torch、torchvision、numpy、Pillow、opencv-python、tqdm、matplotlib 这一套。先建虚拟环境再装别往系统 Python 里灌医学分割项目经常要锁 torch 版本污染了很难回退。# 建议 Python 3.8 ~ 3.10torch 版本按显卡 CUDA 选 python -m venv venv source venv/bin/activate # Windows 用 venv\Scripts\activate pip install -r requirements.txt # 如果 requirements 里没锁 torch手动装匹配 CUDA 的版本 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118装完先确认torch.cuda.is_available()返回 True否则后面训练会默认跑 CPU一个 epoch 能等到你怀疑人生。目录上代码包解压后根目录就是工作目录__pycache__里那堆.pyc是历史编译缓存可以无视但注意里面混着 cpython-37/38/310 多个版本说明这份代码在不同 Python 版本下跑过你本地用哪个版本就以哪个为准别被.pyc干扰。数据目录我一般这样放和代码解耦换数据集不用动代码project/ ├── train.py ├── predict.py ├── model.py ├── dataset.py ├── transforms.py ├── confuse_matrix.py ├── data/ │ ├── train/ │ │ ├── images/ # 原图 │ │ └── masks/ # 标签文件名与 images 一一对应 │ └── val/ │ ├── images/ │ └── masks/2.2 读懂 model.py 里的 U-Net 骨架model.py是整份代码的核心U-Net 的对称结构就在这里。收缩路径是若干「卷积 ReLU 最大池化」的下采样块每下采样一次通道数翻倍、特征图尺寸减半扩展路径是「上采样 拼接对应层特征 卷积」跳跃连接把收缩路径同层的高分辨率特征直接接到扩展路径上。多类别分割的关键改动在最后一层输出通道数等于类别数而不是 1。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_channels3, num_classes4): super().__init__() # 收缩路径 self.down1 DoubleConv(in_channels, 64) self.down2 DoubleConv(64, 128) self.down3 DoubleConv(128, 256) self.down4 DoubleConv(256, 512) self.pool nn.MaxPool2d(2) self.bottleneck DoubleConv(512, 1024) # 扩展路径上采样后与同层特征拼接通道数翻倍再卷积 self.up4 nn.ConvTranspose2d(1024, 512, 2, stride2) self.conv4 DoubleConv(1024, 512) self.up3 nn.ConvTranspose2d(512, 256, 2, stride2) self.conv3 DoubleConv(512, 256) self.up2 nn.ConvTranspose2d(256, 128, 2, stride2) self.conv2 DoubleConv(256, 128) self.up1 nn.ConvTranspose2d(128, 64, 2, stride2) self.conv1 DoubleConv(128, 64) self.out nn.Conv2d(64, num_classes, 1) # 多类别输出通道类别数 def forward(self, x): d1 self.down1(x) d2 self.down2(self.pool(d1)) d3 self.down3(self.pool(d2)) d4 self.down4(self.pool(d3)) b self.bottleneck(self.pool(d4)) u4 self.conv4(torch.cat([self.up4(b), d4], dim1)) u3 self.conv3(torch.cat([self.up3(u4), d3], dim1)) u2 self.conv2(torch.cat([self.up2(u3), d2], dim1)) u1 self.conv1(torch.cat([self.up1(u2), d1], dim1)) return self.out(u1) # 返回 [B, num_classes, H, W] 的 logitsin_channels按输入图像通道给灰度医学影像填 1RGB 填 3。num_classes是分割类别总数二分类任务填 2背景 前景不要填 1否则和交叉熵损失对不上。torch.cat的dim1是通道维拼接这是跳跃连接的本质把上采样的粗特征和同层细特征在通道上叠起来让网络同时看到「这是什么」和「边界在哪」。最后一层用 1x1 卷积把 64 通道压成类别数输出的是未过 softmax 的 logits损失函数里再处理。2.3 第一次前向验证改完模型别急着训练先用随机张量走一遍前向确认输出形状对得上这一步能省掉后面一半的维度报错。import torch from model import UNet device torch.device(cuda if torch.cuda.is_available() else cpu) model UNet(in_channels3, num_classes4).to(device) x torch.randn(2, 3, 256, 256).to(device) # batch2, 3通道, 256x256 with torch.no_grad(): y model(x) print(y.shape) # 期望 torch.Size([2, 4, 256, 256])输出形状是[B, num_classes, H, W]空间尺寸和输入一致这是 U-Net 做像素级分割的前提。如果 H/W 对不上多半是某次池化和上采样次数不匹配或者输入尺寸不是 16 的整数倍——U-Net 下采样 4 次输入边长最好能被 16 整除否则拼接时尺寸差一两个像素就会报错。我一般把输入统一 resize 到 256 或 512省掉这类玄学问题。3. 数据管线与训练循环dataset、transforms 和 train.py 怎么串3.1 dataset.py 与 transforms.py 的配合dataset.py负责把图像和标签读成对transforms.py负责同步增强。医学分割里最容易翻车的地方就是「图像做了随机翻转标签没跟着翻」训练 loss 死活不降。所以增强必须对图像和 mask 用同一组随机参数。import os import numpy as np import torch from torch.utils.data import Dataset from PIL import Image import torchvision.transforms.functional as TF import random class SegDataset(Dataset): def __init__(self, img_dir, mask_dir, img_size256, trainTrue): self.img_dir img_dir self.mask_dir mask_dir self.img_size img_size self.train train self.names sorted(os.listdir(img_dir)) 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(RGB) mask Image.open(os.path.join(self.mask_dir, name)).convert(L) 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) # [3,H,W], 值域 0~1 mask torch.from_numpy(np.array(mask)).long() # [H,W], 类别索引 if self.train and random.random() 0.5: img TF.hflip(img) mask TF.hflip(mask) # 图像翻转标签同步翻转 return img, mask标签 resize 一定要用NEAREST最近邻插值用双线性会把类别索引插成小数比如类别 1 和 2 之间插出 1.5转 long 之后变成莫名其妙的类别这是血泪经验。mask 转成long是因为交叉熵损失要求目标张量是整型类别索引形状[H, W]不是 one-hot。图像转 tensor 后值域是 0~1如果要用 ImageNet 预训练权重还得按均值方差归一化这一步在transforms.py里补。3.2 train.py 的训练循环与损失选择多类别分割的标准损失是交叉熵CrossEntropyLoss内部自带 softmax所以模型输出 logits 直接喂进去不要再手动 softmax。类别不均衡时医学图像里病灶往往只占几个像素加权重或换 Dice 损失。import torch import torch.nn as nn from torch.utils.data import DataLoader from model import UNet from dataset import SegDataset device torch.device(cuda if torch.cuda.is_available() else cpu) model UNet(in_channels3, num_classes4).to(device) train_ds SegDataset(data/train/images, data/train/masks, trainTrue) loader DataLoader(train_ds, batch_size4, shuffleTrue, num_workers4) # 类别权重背景多就压低背景权重病灶少就抬高 weights torch.tensor([0.2, 1.0, 1.0, 1.0]).to(device) criterion nn.CrossEntropyLoss(weightweights) optimizer torch.optim.Adam(model.parameters(), lr1e-3) for epoch in range(50): model.train() total_loss 0 for img, mask in loader: img, mask img.to(device), mask.to(device) optimizer.zero_grad() out model(img) # [B, C, H, W] loss criterion(out, mask) # mask: [B, H, W] long loss.backward() optimizer.step() total_loss loss.item() print(fepoch {epoch}, loss {total_loss/len(loader):.4f})batch_size受显存限制256x256 输入、4 类输出8G 显存大概能跑 batch 4~8。lr用 1e-3 起步loss 震荡就降到 1e-4。num_workers在 Windows 上设 0 更稳设大了容易卡在数据加载。训练时盯着 loss 曲线如果前几个 epoch 就掉到很低但验证集一塌糊涂八成是标签和图像没对齐或者标签类别索引从 1 开始而损失期望从 0 开始。3.3 验证与指标confuse_matrix.py 怎么用confuse_matrix.py是这份包里容易被忽略但很实用的模块它算的是混淆矩阵进而能推出每个类别的 IoU 和 Dice。分割任务光看 loss 不够loss 低不代表边界分得好。import numpy as np def compute_confusion(preds, targets, num_classes): # preds/targets: 展平后的整型数组 mask (targets 0) (targets num_classes) hist np.bincount( num_classes * targets[mask].astype(int) preds[mask], minlengthnum_classes ** 2 ).reshape(num_classes, num_classes) return hist def iou_from_confusion(hist): # 对角线是预测正确的像素 inter np.diag(hist) union hist.sum(axis1) hist.sum(axis0) - inter return inter / np.maximum(union, 1) # 每类 IoUhist的行是真实类别、列是预测类别对角线越大越好。iou_from_confusion返回每个类别的 IoU背景类通常虚高重点看病灶类的值。验证时把模型输出argmax(dim1)得到预测类别图和 mask 一起展平送进compute_confusion。如果某个类别 IoU 长期为 0先查这个类别的像素在训练集里是不是几乎没有再查标签里这个类别的索引有没有被 resize 破坏。4. 多类别分割的配置与常见翻车排查4.1 多类别与二分类的配置差异同一份代码二分类和多类别的差别集中在三个地方改错一个就报错或静默出错。下面这张表是我实际调的时候总结的对照配置项二分类多类别模型输出通道num_classes2背景前景类别总数 N标签格式0/1 整型索引0~N-1 整型索引损失函数CrossEntropyLoss 或 BCECrossEntropyLoss预测取类别argmax(dim1)argmax(dim1)标签 resize 插值NEARESTNEAREST注意标签索引必须从 0 连续到 N-1。有些标注工具导出的 mask 像素值是 0、128、255直接喂进去会报「target out of bounds」得先做一次映射把 128 映射成 1、255 映射成 2。import numpy as np # 把 0/128/255 的标注映射成 0/1/2 def remap_mask(mask): mapping {0: 0, 128: 1, 255: 2} out np.zeros_like(mask, dtypenp.int64) for k, v in mapping.items(): out[mask k] v return out4.2 避坑与排查五个真实翻车记录现象一训练 loss 一直不降或者降到某个值就卡住。原因图像和标签增强不同步翻转/旋转只作用在图像上网络学的是错位对应关系。 解决所有随机增强必须对 img 和 mask 用同一随机种子或同一判断分支像 3.1 里那样if random.random() 0.5同时翻转两者。现象二报错Expected target size [B, H, W], got [B, H, W, C]。原因标签被做成了 one-hot形状多了一维而CrossEntropyLoss要的是类别索引。 解决dataset 里 mask 保持[H, W]的 long 张量不要 to_onehot如果上游给的是 one-hot用argmax(dim-1)压回去。现象三验证集 IoU 正常但 predict.py 出图全黑或全是一类。原因预测时忘了对输出做argmax直接把 logits 当类别图保存或者保存时没做归一化。 解决pred torch.argmax(out, dim1)得到[B, H, W]再乘一个缩放系数比如 255//num_classes存成灰度图方便肉眼检查。现象四显存溢出batch_size 降到 1 还爆。原因输入分辨率太大或者上采样用了ConvTranspose2d且通道没控制好。 解决先把输入降到 256或者把num_workers调小、开torch.cuda.amp混合精度实在不行把 bottleneck 的 1024 通道砍到 512医学小数据集用不着那么宽。现象五某个类别 IoU 恒为 0。原因该类别像素在训练集占比极低被背景淹没或者标签映射时这个类别的值没被正确重映射。 解决给CrossEntropyLoss加类别权重或改用 Dice CE 组合损失同时统计训练集每类像素占比占比低于千分之一的类别要考虑过采样。5. 从能跑到好用预测、可视化与一个提分技巧训练跑通只是起点真正决定这份代码包好不好用的是预测和验证环节。predict.py的职责是加载权重、对单张或批量图像推理、把argmax后的类别图存下来。我习惯在预测里加一段叠加可视化把原图和分割结果半透明叠一起边界对不对一眼就能看出来比盯着 IoU 数字直观得多。import torch import numpy as np from PIL import Image from model import UNet import torchvision.transforms.functional as TF device torch.device(cuda if torch.cuda.is_available() else cpu) model UNet(in_channels3, num_classes4).to(device) model.load_state_dict(torch.load(best.pth, map_locationdevice)) model.eval() img Image.open(data/val/images/sample.png).convert(RGB) x TF.to_tensor(TF.resize(img, [256, 256])).unsqueeze(0).to(device) with torch.no_grad(): out model(x) pred torch.argmax(out, dim1).squeeze(0).cpu().numpy() # [256,256] # 叠加可视化原图 70% 分割色 30% color_map np.array([[0,0,0],[255,0,0],[0,255,0],[0,0,255]], dtypenp.uint8) overlay color_map[pred] base np.array(TF.resize(img, [256, 256])) blend (base * 0.7 overlay * 0.3).astype(np.uint8) Image.fromarray(blend).save(overlay.png)color_map的行数等于类别数每行一个 RGB 颜色color_map[pred]把[H, W]的类别索引直接映射成[H, W, 3]彩色图这一步比逐像素循环快几个数量级。load_state_dict的map_location在 CPU 推理时必加否则加载 GPU 保存的权重会报设备不匹配。一个提分技巧医学分割里边界像素最容易错可以在损失里对边界加权。做法是先对 mask 做形态学腐蚀和膨胀两者相减得到边界带给边界带上的像素更高权重。常见做法是用scipy.ndimage的binary_erosion和binary_dilation生成边界掩码再乘进损失权重图。这个改动不大但在病灶边界模糊的数据集上Dice 通常能涨两三个点。我一般会在正式训练前先用小学习率跑 5 个 epoch看边界类的 IoU 有没有提升没提升就撤掉别硬堆。从那以后我每次拿到一份分割代码包都强制先走一遍「随机张量前向 → 单 batch 过拟合 → 全量训练」三步单 batch 过拟合能在一分钟内暴露标签错位、类别越界、损失配置这些最坑的问题比直接开训省下大半天。希望这份拆解帮到你把这份 U-Net 代码包真正跑成自己数据集上的基线。本文还有配套的精品资源点击获取
返回列表