ARTICLE DETAIL

资讯详情

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

PyTorch实现U-Net图像分割:从数据加载到训练评估全解析

PyTorch实现U-Net图像分割:从数据加载到训练评估全解析 图像分割是计算机视觉里一个很典型的任务。它和分类、检测不同分类判断整张图属于哪一类检测框出目标的位置而分割要精确到像素级别把图像里的每个像素都标成对应的类别。用 PyTorch 实现 U-Net 架构做图像分割是很多入门者学习像素级预测任务最直接的一条路径也是医学影像、卫星图像、工业质检等场景里经常出现的方案。这篇文章会从环境准备、网络结构、数据加载、训练评估到问题排查按实际落地顺序拆一遍。我默认读者已经会跑简单的 PyTorch 分类模型。如果你连 PyTorch 都还没装上先看第 2 节如果你已经跑通过分类模型只是第一次接触分割任务可以直接跳到第 3 节看网络结构。整篇文章不会引入复杂的模型变体只关注一个目标把 U-Net 在 PyTorch 里跑通并且知道每一步调参与排查的逻辑是什么。1. 图像分割和 U-Net先明确这套架构解决什么问题1.1 分割任务和分类任务的核心差异分类模型最后输出的是一个向量向量长度等于类别数softmax 之后取最大下标就是预测类别。分割模型不一样它要输出一张和输入分辨率接近的图图上的每个像素都有一个类别标签。这个差异决定了整个训练流程的走向数据组织方式不同。分类只需要图像和标号分割需要图像和对应的掩码图。损失函数计算方式不同。分类对整图算一个交叉熵分割是对每个像素分别算损失再取平均。评估指标不同。分类看 accuracy分割更看重 IoU、Dice 这类像素级重叠指标。网络输出头不同。分割网络的最后一般是一个 1x1 卷积把通道数映射成类别数再接 sigmoid 或 softmax。很多人刚接触分割时习惯拿分类的思维去套结果在数据加载和 loss 计算上反复出问题。先把这个本质差异记住后面会少踩很多坑。1.2 U-Net 解决的实际问题U-Net 最早出现在医学图像分割领域论文名字里的 U 是因为网络结构画出来像一个 U 形。它解决的核心问题是既要有全局语义信息来判断目标是什么又要有局部细节信息来判断目标的边界在哪里。普通卷积网络在层层下采样之后语义信息越来越强但空间细节越来越弱。如果直接把最后一层特征图上采样回原尺寸边界会是糊的。U-Net 通过跳跃连接把编码器每一层的高分辨率特征直接拼到解码器对应层让解码器在恢复分辨率的同时还能拿到底层细节。这套思路到今天仍然有效。很多新的分割模型包括一些基于 Transformer 的结构内部也在用类似的多尺度特征融合思想。所以 U-Net 不只是老模型它是理解分割任务的一个很好的骨架。1.3 典型应用场景U-Net 最常见的应用是医学图像分割比如 CT 影像里的器官分割、MRI 图像里的病灶分割、细胞显微图像里的细胞核分割。这类场景有个特点目标边缘复杂标注成本高样本量往往不大。U-Net 在少量数据上也能训练出可用的结果这正好契合需求。除此之外U-Net 也常用于遥感图像里的建筑物、道路、水体分割。工业质检中的缺陷区域分割。自动驾驶场景中的路面、车辆、行人分割。人像分割、图像抠图等任务。本文的代码示例以二分类分割为主也就是每个像素只区分前景和背景。如果你要做多类别分割改动集中在 loss、输出通道数和评估逻辑上网络骨架不用换。2. 运行环境准备从安装到项目目录一次理清2.1 安装 PyTorch 的常见方式和版本选择每次写 PyTorch 实战环境都是第一个门槛。最常见的安装方式是用 Anaconda 创建独立环境然后用 pip 安装 PyTorch。不建议直接把 PyTorch 装到 base 环境里因为不同项目依赖的 PyTorch 版本可能不一样隔离环境能减少很多冲突。安装命令要根据你的 CUDA 版本来选择。如果你用 GPU 跑先确认显卡驱动支持什么 CUDA 版本再选择对应的 PyTorch 版本。教程里经常出现的一类报错是warning: you need pytorch with cu130 or higher to use optimized cuda operations这类提示说明当前 PyTorch 版本和 CUDA 版本不匹配或者当前 PyTorch 缺少针对你所用 GPU 架构的算子优化。遇到这种提示先不要急着换显卡去官方安装页面重新选择匹配的版本。确认安装是否成功的标准操作是这样python -c import torch; print(torch.__version__); print(torch.cuda.is_available())如果torch.cuda.is_available()返回True说明 GPU 可用。注意这里只代表 PyTorch 能找到你的 CUDA 环境不代表跑起来一定不会出其他问题。我一般还会再跑一句print(torch.cuda.get_device_name(0))确认真的读到了正确的显卡。如果你的环境只能用 CPU也不用放弃。U-Net 在小尺寸图片上 CPU 也能跑只是训练时间会长一些。把图片尺寸控制在 128x128 或 256x256把 epoch 数调小一点先把流程跑通后面再考虑 GPU。2.2 硬件条件怎么判断运行环境准备阶段我建议按这个优先级确认显存、内存、磁盘、CPU 核心数。显存决定能开多大的 batch size。U-Net 在 256x256 输入下batch size 设为 8 一般需要 6GB 以上显存。如果你的显卡只有 4GB就降成 4 或者 2。内存影响数据加载。如果数据集很大建议在 DataLoader 里设置合适的num_workers默认 0 也能跑但数据预处理会拖慢训练。磁盘影响数据集读取速度。图像分割数据集通常是一堆 PNG 文件磁盘 IO 慢会直接拉低训练速度。如果显存不够第一反应不要是换显卡先检查这几个地方batch size 是不是开太大、输入分辨率是不是可以降、num_workers是不是设置过高导致内存占用上涨。很多时候一个参数就能解决问题。注意低配机器能跑通但不代表能按默认参数稳定跑完整个训练。显存不足、内存暴涨、训练中断这类问题在低配环境里尤其常见先把 batch size 和图片尺寸降下来再试。2.3 项目目录与数据集组织一个清晰的项目目录能省掉大量调试时间。我一般会这样组织unet_segmentation/ ├── data/ │ ├── images/ │ └── masks/ ├── models/ │ └── unet.py ├── utils/ │ ├── dataset.py │ ├── metrics.py │ └── visualize.py ├── checkpoints/ ├── train.py └── predict.py目录结构本身没有标准答案核心原则是训练代码、模型定义、数据加载、输出目录分开。这样后面跑批量实验时换数据集、换模型、换超参数都只改一个小文件不用在长脚本里找关键行。官方或开源数据集的组织方式通常分两种一种是 images 和 masks 两个独立文件夹文件名一一对应另一种是单个文件夹下每个样本有子文件夹里面同时放原图和掩码。无论哪种写 Dataset 时都建议提前打印一两条样本的路径和尺寸做验证不要直接开训练。3. U-Net 网络结构拆解编码器、跳跃连接、解码器3.1 基本卷积块U-Net 的底层单元是连续两次 3x3 卷积中间夹一个 ReLU 激活有些实现还会加 BatchNorm。这个结构在编码器和解码器里都会被反复用。用 PyTorch 可以写成这样import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size3, padding1), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue), nn.Conv2d(out_channels, out_channels, kernel_size3, padding1), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue) ) def forward(self, x): return self.conv(x)为什么要连续两层卷积而不是一层两层 3x3 卷积堆叠感受野相当于一层 5x5 卷积但参数量更少而且中间多了一次非线性变换表达能力更强。这是从 VGG 时代就验证过的经验U-Net 沿用了这个设计。3.2 编码器逐层下采样提取语义特征编码器的任务是把输入图像逐步缩小并增加通道数。每一步结构是先过 DoubleConv 提取特征再用nn.MaxPool2d(2)把空间尺寸减半。从 1 通道输入开始经过四到五次下采样特征图会从 572x572 一路缩到 32x32通道数从 1 涨到 512 或 1024。通道数增加的意义在于空间分辨率下降之后每个像素需要承载更多的语义信息。下采样也带来了平移不变性让模型对目标位置的微小变化不那么敏感。代价是边界细节损失这正是后面解码器和跳跃连接要补偿的。编码器一个常见的问题是过度下采样。如果输入图像本身不大比如 96x96做五次下采样之后只剩 3x3信息几乎丢没了。所以结构里的下采样层数要根据输入分辨率来定不能照抄原始论文的四层五层。3.3 跳跃连接U-Net 最核心的设计跳跃连接就是把编码器的某一层输出通过torch.cat拼接到解码器的对应层输入上。原始 U-Net 的跳跃连接直接把编码器特征拷贝过去后来也有注意力门控版本、把连接换成加法融合的版本但最经典的实现就是通道拼接。为什么拼接比单纯相加更好拼接保留了编码器特征的全部通道信息解码器在训练时可以通过卷积自动学习如何融合语义特征和细节特征。相加是一种强约束强制两个特征在数值上对齐拼接则给网络更大的自由度虽然参数量和内存占用会高一些但效果通常更稳。实现时要注意通道数匹配。比如编码器第三层输出是 256 通道解码器对应层的当前特征也是 256 通道拼接后就变成 512 通道下一层 DoubleConv 的输入通道必须是 512。这个数字很容易写错建议写代码时先画一张图表把每一层的输入输出通道数列出来。3.4 解码器恢复分辨率并整合特征解码器的每一步包含两个操作先用转置卷积或双线性插值把特征图放大一倍再和对应编码器层的跳跃连接输出拼接。放大方式有讲究转置卷积可以学习上采样参数但容易出现棋盘格伪影需要调好 kernel size 和 stride。双线性插值没有可学习参数结果更平滑配合后面的卷积层也能学到调整能力。原始 U-Net 用的是转置卷积。现代实现里很多人改成双线性插值加 1x1 卷积稳定性更好一些。如果你刚入门建议先用转置卷积因为它在很多公开代码里最常见理解之后再换。关键参数是kernel_size2, stride2这样特征图尺寸刚好翻倍。解码器的输出要经过最后一层 1x1 卷积把通道数映射成类别数。二分类任务输出 1 个通道多分类任务输出类别数个通道。这里输出的 logits 不需要自己加 sigmoid因为 PyTorch 的BCEWithLogitsLoss和CrossEntropyLoss会在内部处理数值稳定的激活计算。3.5 完整网络组装一个适合小数据集的轻量 U-Net 完整定义如下class UNet(nn.Module): def __init__(self, in_channels1, num_classes1, features[64, 128, 256, 512]): super().__init__() self.encoder nn.ModuleList() self.pool nn.MaxPool2d(2) for f in features: self.encoder.append(DoubleConv(in_channels, f)) in_channels f self.bottleneck DoubleConv(features[-1], features[-1] * 2) self.up nn.ModuleList() self.decoder nn.ModuleList() self.upsample nn.Upsample(scale_factor2, modebilinear, align_cornersTrue) for f in reversed(features): self.up.append( nn.ConvTranspose2d(f * 2, f, kernel_size2, stride2) ) self.decoder.append(DoubleConv(f * 2, f)) self.final_conv nn.Conv2d(features[0], num_classes, kernel_size1) def forward(self, x): skip_connections [] for enc in self.encoder: x enc(x) skip_connections.append(x) x self.pool(x) x self.bottleneck(x) skip_connections skip_connections[::-1] for i in range(len(self.decoder)): x self.up[i](x) x torch.cat([x, skip_connections[i]], dim1) x self.decoder[i](x) return self.final_conv(x)这个版本的通道数设计比原始论文小很多适合 256x256 以下的输入图片。原始 U-Net 的编码器是 64、128、256、512、1024对小数据集来说参数量偏大容易过拟合。4. 数据加载与预处理分割任务最容易出错的地方4.1 图像和掩码的对应关系分割任务的训练数据是一对一关系每一张输入图像对应一张掩码图。掩码图中每个像素的灰度值就是该像素的类别标签。二分类场景里背景像素通常是 0前景像素是 255 或 1需要把 255 归一化成 1。写数据加载代码时最容易犯的错误是图像和掩码没有对齐。比如文件名排序不一致、某个掩码缺失、输入图像是 3 通道彩色图而掩码是单通道图这些都会导致训练时模型学到错误对应关系。我建议数据加载代码里加上形状校验assert img.size mask.size, fimage and mask size mismatch: {img.size} vs {mask.size} assert len(mask.shape) 2, fmask should be single channel, got shape {mask.shape}这个断言在训练前跑一遍能拦截掉大量无声的 bug。4.2 归一化和增强策略输入图像一般做标准化处理。灰度图用 mean0.5, std0.5 把像素归一化到 [-1, 1]彩色图可以用每个通道独立的均值和标准差也可以简单用 0.5。掩码图不要做同样的标准化掩码是标签必须保持 0/1 或者 0/1/2 这种类别值归一化会把标签搞坏。数据增强方面分割任务和分类任务有个关键差异对图像做的几何变换必须同步应用到掩码上。翻转、旋转、裁剪这类变换可以同时作用于两者但颜色抖动、亮度变化只能作用于输入图像因为标签不依赖颜色信息。PyTorch 里处理这类同步变换常见做法是写一个自定义的__getitem__在返回前手动做变换。torchvision 的v2版本也提供了同时处理图像和掩码的接口如果你用的是最新版可以试试逻辑上更省事。4.3 Dataset 的典型写法一个基础的 Dataset 实现如下import os from PIL import Image import torch from torch.utils.data import Dataset from torchvision import transforms class SegmentationDataset(Dataset): def __init__(self, image_dir, mask_dir, image_size(256, 256)): self.image_dir image_dir self.mask_dir mask_dir self.image_size image_size self.image_files sorted(os.listdir(image_dir)) self.mask_files sorted(os.listdir(mask_dir)) self.image_transform transforms.Compose([ transforms.Resize(image_size), transforms.ToTensor(), transforms.Normalize(mean[0.5], std[0.5]) ]) self.mask_transform transforms.Compose([ transforms.Resize(image_size, interpolationImage.NEAREST), transforms.ToTensor() ]) def __len__(self): return len(self.image_files) def __getitem__(self, idx): img_path os.path.join(self.image_dir, self.image_files[idx]) mask_path os.path.join(self.mask_dir, self.mask_files[idx]) image Image.open(img_path).convert(L) mask Image.open(mask_path).convert(L) image self.image_transform(image) mask self.mask_transform(mask) # 把掩码从 0-1 范围转成长整型标签 mask (mask 0.5).long() return image, mask掩码用NEAREST插值缩放这个细节很重要。如果用双线性插值掩码边界会变成模糊的过渡灰阶导致标签出现不存在的中间值。用最近邻插值能保证标签值不被打散。如果你处理的掩码是 PNG 格式读取时颜色模式要保持一致。PIL 读取灰度 PNG 时加上.convert(L)读取 RGB PNG 时加.convert(RGB)避免单通道和三通道数据源混用导致张量形状错乱。5. 训练流程损失函数、优化器和关键参数5.1 二分类分割用 BCEWithLogitsLoss多分类用 CrossEntropyLoss分割任务最常用的两个损失函数二分类BCEWithLogitsLoss网络输出 shape 为(N, 1, H, W)掩码 shape 为(N, H, W)的长整型。多分类CrossEntropyLoss网络输出 shape 为(N, C, H, W)掩码 shape 为(N, H, W)其中 C 是类别数。很多分割任务还会在交叉熵基础上叠加 Dice Loss 或 IoU Loss。这是因为像素分类存在严重的类别不平衡问题比如医学图像里病灶通常只占整张图的很小比例模型很容易学会把所有像素预测成背景loss 依然很低。一个简单的 Dice Loss 实现def dice_loss(pred, target, smooth1.0): pred torch.sigmoid(pred) pred pred.view(pred.size(0), -1) target target.view(target.size(0), -1).float() intersection (pred * target).sum(dim1) union pred.sum(dim1) target.sum(dim1) dice (2.0 * intersection smooth) / (union smooth) return 1.0 - dice.mean()实际使用中一般把 BCE 和 Dice 按比例相加比如total_loss bce_loss dice_loss。两个损失的数值量级可能差很多如果你发现训练早期 loss 波动很大可以先固定权重比如 BCE 占 0.7、Dice 占 0.3再根据结果调整。5.2 优化器、学习率和批次大小优化器建议从 Adam 开始。默认学习率1e-3在大多数分割任务下能跑出一个可接受的结果但更稳妥的做法是从1e-4开始因为分割任务的损失曲面更复杂学习率太大容易在早期震荡。如果我训练时 loss 下降速度太慢第一步不是把学习率调大 10 倍而是打印出来看看学习率有没有被 scheduler 降到过低或者输入数据是否真的加载正确。很多时候 loss 不降是因为标签全为 0 或者数据没有归一化跟学习率没有关系。batch size 的设置可以直接参考显存。256x256 的输入8GB 显存建议 batch size 设为 4 到 816GB 显存可以到 16。显存不够时优先降 batch size不要轻易改模型通道数因为改结构会引入新的调试变量。5.3 训练循环和日志观察训练循环不需要复杂的封装。以 PyTorch 标准写法为例from torch.utils.data import DataLoader from torch.optim import Adam from torch.nn import BCEWithLogitsLoss dataset SegmentationDataset(data/images, data/masks) dataloader DataLoader(dataset, batch_size4, shuffleTrue, num_workers2) model UNet(in_channels1, num_classes1) optimizer Adam(model.parameters(), lr1e-4) criterion BCEWithLogitsLoss() device torch.device(cuda if torch.cuda.is_available() else cpu) model.to(device) num_epochs 50 for epoch in range(num_epochs): model.train() running_loss 0.0 for images, masks in dataloader: images images.to(device) masks masks.to(device).float() outputs model(images) loss criterion(outputs, masks) optimizer.zero_grad() loss.backward() optimizer.step() running_loss loss.item() * images.size(0) epoch_loss running_loss / len(dataset) print(fEpoch {epoch1}/{num_epochs}, Loss: {epoch_loss:.4f}) if (epoch 1) % 10 0: torch.save(model.state_dict(), fcheckpoints/unet_epoch_{epoch1}.pth)训练中要盯的指标不只是 loss还包括显存占用、单 epoch 耗时和数据加载速度。如果单 epoch 耗时异常长先看num_workers是否为 0、磁盘是否慢、图像尺寸是否过大。保存模型时不要只存最后一轮的权重。分割任务在训练后期 IoU 可能不再提升但 weight 仍然有价值。我建议每 10 到 20 个 epoch 保存一次检查点并且把当前验证集的指标写进日志后面可以挑最优的权重来做推理。6. 推理、评估与可视化不能只看 loss6.1 推理时的输出后处理训练完成后推理阶段要把模型输出转成可视化结果。模型输出的是 logits二分类时每个像素有一个实数值需要经过 sigmoid 再阈值化with torch.no_grad(): model.eval() output model(image_tensor) # shape: (1, 1, H, W) prob torch.sigmoid(output) # shape: (1, 1, H, W) mask (prob 0.5).float() # shape: (1, 1, H, W)阈值 0.5 是默认选择但不是所有场景都适用。如果预测出来的前景区域偏大可以调高阈值比如 0.6 或 0.7如果偏小就调低。这个现象在医学图像里很常见因为病灶边界模糊模型输出的概率往往是渐变的。更合理的做法是看一眼概率图的直方图分布。如果大部分像素的概率都在 0.9 以上或 0.1 以下说明模型置信度很高阈值影响不大如果大量像素集中在 0.4 到 0.6 之间说明边界区域犹豫阈值需要结合具体业务场景调整。6.2 评估指标IoU 和 Dice分割任务不能只看 loss因为像素准确率在背景占绝大多数的情况下会虚高。一张图 95% 是背景模型把所有像素都预测成背景像素准确率也有 95%但这个模型没有任何实际价值。常用指标有两个IoU预测结果和真实掩码的交集面积除以并集面积。Dice两倍的交集面积除以两者面积之和。两者数值相关Dice 通常会比 IoU 高一些。计算时要注意两个细节第一是sigmoid(pred) 0.5之后再算指标不要用 logits 直接算第二是一次性把整个验证集的预测结果累加再除以总样本数而不是对每个样本单独算完再取平均后者会更偏重大图中表现好的样本。6.3 可视化预测结果训练结束之后强烈建议把几张验证图的预测掩码、真实掩码和原图并排保存成图片。可视化能让你直接从视觉上判断模型的问题是边界不光滑是小目标漏检还是前后景粘连。一段简单的可视化逻辑import matplotlib.pyplot as plt plt.figure(figsize(12, 4)) plt.subplot(1, 3, 1) plt.imshow(image.squeeze(), cmapgray) plt.title(Input) plt.subplot(1, 3, 2) plt.imshow(mask.squeeze(), cmapgray) plt.title(Ground Truth) plt.subplot(1, 3, 3) plt.imshow(pred.squeeze(), cmapgray) plt.title(Prediction) plt.savefig(result_sample.png)可视化这一步不能省。数值指标可以告诉你模型好不好但不能直接告诉你哪里不好。看到预测图上的具体错误模式你才能决定下一步是增加数据增强、调整损失函数权重还是修改网络结构。7. 常见问题与排查顺序7.1 输入输出尺寸不匹配这是 U-Net 代码里最常见的报错。原因通常是下采样和上采样的次数不同或者跳跃连接拼接时通道数不一致。排查顺序打印网络输入输出尺寸确认模型定义阶段能跑通一次前向传播。把输入设为固定尺寸比如 256x256不要在一开始就支持任意尺寸。检查 MaxPool 和 Upsample 的层数是否严格对应。检查跳跃连接时两个张量的 H、W 是否一致。如果输入尺寸不是 2 的整数次幂下采样后可能出现奇数尺寸导致拼接时报错。最稳妥的做法是用一个随机张量初始化输入在进入训练前先跑一次 forward验证尺寸没问题再开始。7.2 训练 loss 不下降先确认数据本身。打印一张图像和掩码检查掩码标签分布是否合理、图像像素值范围是否正常。如果标签全是 0模型只会学会输出全零预测loss 可能很低但没有意义。再确认模型是否真的在训练。用一个小 batch 跑几步打印每一步的 loss看它在优化器 step 之后有没有变化。如果 loss 完全没有波动多半是requires_grad的问题或者模型输入输出没有连接到损失函数上。然后检查学习率。Adam 在1e-4到3e-4之间通常能正常收敛。如果 loss 一直在高位震荡可以试着把学习率降到1e-5再观察几十个 epoch。最后确认归一化。输入图像和掩码都需要符合 loss 函数的预期BCEWithLogitsLoss 需要输入 logits但掩码要保持在 0 到 1 的范围CrossEntropyLoss 的掩码必须是长整型取值为 0 到 C-1。7.3 预测结果全黑或全白预测全黑可能原因有三个sigmoid 阈值设置过高模型输出概率普遍偏低。训练数据里前景像素本来就很少模型学到偏向背景的输出。模型没有训练好最后一层权重未收敛。预测全白的情况较少通常是因为测试数据分布和训练数据差异太大或者推理时输入的归一化方式跟训练时不一致。还有一种情况是掩码加载时没有转成正确的类别值。如果掩码图里前景值是 255但没有归一化到 1模型会把它当多分类任务处理输出结果会变得很奇怪。7.4 环境相关报错环境类问题在前 30 分钟的排查量通常大于模型问题。CUDA 版本不匹配统一在官网按 CUDA 版本重新安装 PyTorch。显存不足降低 batch size、降低分辨率、减少num_workers。依赖版本冲突用独立 conda 环境不要混装 torchvision 和 PyTorch 的不同版本。数据加载卡住先设num_workers0跑一次排除多进程读取的锁和异常问题。这些问题的共性排查顺序是先看完整报错栈再看具体报错类型最后才去改代码。不要看到英文报错就直接搜解决方案先把变量的值打印出来很多问题自己就能定位。结尾跑通一个 U-Net 图像分割项目真正关键的其实不是网络结构写了多少行而是你能不能把数据、训练、评估这条链路完整打通。很多初学者卡在环境安装或者卡在掩码预处理又或者训练了很久发现 loss 降不下去最后才发现是标签数据没有对齐。这些坑在这条任务链上几乎是必经的。我个人更建议先把单张图片的训练跑通用很小的 batch size 和很小的 epoch 数验证整个流程再去做数据增强、调整损失函数、上大分辨率。每一步只引入一个变化出了问题也容易定位。等模型在验证集上有了稳定表现再去考虑批量预测、接口封装和部署那又是另一个阶段的事情了。
返回列表