PyTorch U-Net语义分割实战:从训练到推理的完整代码解析

发布时间:2026/10/2 10:01:50
PyTorch U-Net语义分割实战:从训练到推理的完整代码解析 简介这份资源是面向深度学习初学者与高校学生的PyTorch图像语义分割实战代码包以经典U-Net网络为核心可用于课程设计、期末大作业或入门级科研复现。压缩包共18个文件约2.15MB包含7个Python脚本、3个XML配置、2个pyc缓存、1个pth权重、1个jpg与1个png示例图、1个md说明文档等覆盖训练、测试、数据集加载与网络定义等完整流程。已有958人学习下载说明其作为教学案例具备一定参考价值。代码注释较为充分新手可据此理解U-Net的编码器-解码器结构、跳跃连接与逐像素分类原理并直接运行train.py与test.py完成训练和推理配套的end.pth权重与示例图片便于快速验证效果README则提供环境与使用说明。整体目录清晰、部署门槛低适合需要一份可运行、可修改的语义分割基线方案来支撑作业或课程项目的读者。1. 拆开这份 PyTorch U-Net 语义分割包它到底能跑出什么结果如果你手头正压着一个图像语义分割的期末大作业或者想找一个能直接跑通、结构清晰的 U-Net 训练与测试代码这份基于 PyTorch 的 U-Net 图像语义分割训练和测试代码包大概率能省掉你从零搭网络、写数据加载、调损失函数的那几天。它解决的不是“什么是语义分割”这种概念问题而是把编码器-解码器结构、跳跃连接、逐像素分类、训练循环、验证指标、模型保存与推理测试这一整条链路用可运行的 Python 源码串了起来。适合谁适合已经装好 PyTorch 环境、知道张量是什么、但第一次完整做分割任务的人也适合熟手拿它当 baseline快速替换数据集和骨干网络。下面我按“资源是什么 → 怎么用 → 坑在哪”的顺序把这份代码包拆开讲透。2. U-Net 结构与数据管道为什么这样搭每一步在算什么2.1 编码器-解码器与跳跃连接的真实作用U-Net 的核心不是“深”而是“对称”。编码器每经过一次卷积加下采样特征图尺寸减半、通道数翻倍语义信息越来越强但空间位置越来越粗。解码器做上采样把尺寸逐步恢复可如果只靠上采样边缘和细小目标会糊掉。跳跃连接把编码器同尺度的特征直接拼到解码器对应层让浅层的纹理、边界信息参与最终逐像素判断。这份代码包里收缩路径通常是四次下采样扩张路径对应四次上采样最后接一个 1×1 卷积把通道压到类别数。常见做法是每个卷积块用两个 3×3 卷积加 ReLU再跟一个 2×2 最大池化解码器用转置卷积或双线性插值加卷积。选哪种转置卷积能学参数但容易产生棋盘伪影双线性插值更稳我一般先用插值跑通再换转置卷积对比。2.2 数据加载与标签对齐分割任务最容易翻车的地方分类任务标签是一个数分割任务标签是一张和原图同尺寸的掩码图。代码包里一般会有一个继承Dataset的类在__getitem__里同时读原图和掩码做同样的缩放或裁剪再转成张量。这里必须保证图像和掩码的几何变换完全同步否则训练 loss 会降但预测全是乱码。下面这段是常见写法的骨架我按可复现的方式补全import os import torch from torch.utils.data import Dataset, DataLoader from PIL import Image import torchvision.transforms as T import torchvision.transforms.functional as F class SegDataset(Dataset): def __init__(self, img_dir, mask_dir, size(256, 256)): self.img_dir img_dir self.mask_dir mask_dir self.size size 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 F.resize(img, self.size) mask F.resize(mask, self.size, interpolationT.InterpolationMode.NEAREST) img F.to_tensor(img) mask torch.from_numpy(np.array(mask)).long() return img, mask loader DataLoader(SegDataset(data/images, data/masks), batch_size4, shuffleTrue, num_workers2)逻辑说明convert(L)把掩码转成单通道灰度每个像素值就是类别索引interpolationNEAREST是关键掩码不能用双线性插值否则会插出 0.5 这种不存在的类别。参数上batch_size受显存限制256×256 输入下 4 到 8 比较稳num_workers在 Windows 上建议设 0 或 2设大了容易卡在共享内存。图像归一化可以用 ImageNet 均值方差也可以只用 0.5看你的数据分布。2.3 损失函数与评估指标别只看 accuracy分割任务里背景像素往往占 80% 以上如果只报像素准确率模型全预测背景也能拿高分这就是典型的“看起来能跑实际没用”。代码包里一般会用交叉熵进阶一点用 Dice Loss 或交叉熵加 Dice 的组合。交叉熵管每个像素的分类置信度Dice 管预测区域和真实区域的重叠度两者互补。评估时至少看 IoU 或 Dice 系数按类别算再平均。下面是一个组合损失的写法import torch.nn as nn import torch.nn.functional as F class DiceLoss(nn.Module): def __init__(self, smooth1e-6): super().__init__() self.smooth smooth def forward(self, logits, targets): probs F.softmax(logits, dim1) targets_onehot F.one_hot(targets, num_classeslogits.shape[1]) targets_onehot targets_onehot.permute(0, 3, 1, 2).float() intersection (probs * targets_onehot).sum(dim(2, 3)) union probs.sum(dim(2, 3)) targets_onehot.sum(dim(2, 3)) dice (2 * intersection self.smooth) / (union self.smooth) return 1 - dice.mean() criterion nn.CrossEntropyLoss() DiceLoss()参数说明smooth防止分母为零one_hot把掩码转成和 logits 同通道数的独热形式dim(2,3)是在空间维度求和。注意 Dice Loss 对类别不平衡更敏感如果某些类别样本极少可以给交叉熵加类别权重。训练时先跑几个 epoch 看 loss 是否稳定下降如果 Dice 一直不动先检查掩码像素值是不是从 0 开始连续很多数据集背景是 255直接送进CrossEntropyLoss会报越界或算出错误结果。3. 训练与测试全流程从命令行到推理输出3.1 训练脚本的关键参数与断点续训训练脚本一般包含模型实例化、优化器、学习率调度、epoch 循环、验证、保存最优权重。优化器常用 Adam 或 SGDAdam 起步快SGD 调好了泛化更好。学习率初始 1e-3 到 1e-4配合ReduceLROnPlateau在验证 IoU 不升时降学习率。下面是一个可抄的训练循环骨架import torch from torch.optim import Adam from torch.optim.lr_scheduler import ReduceLROnPlateau device torch.device(cuda if torch.cuda.is_available() else cpu) model UNet(in_channels3, num_classes2).to(device) optimizer Adam(model.parameters(), lr1e-3) scheduler ReduceLROnPlateau(optimizer, modemax, patience5) best_iou 0.0 for epoch in range(50): model.train() for imgs, masks in train_loader: imgs, masks imgs.to(device), masks.to(device) optimizer.zero_grad() logits model(imgs) loss criterion(logits, masks) loss.backward() optimizer.step() model.eval() iou evaluate(model, val_loader, device) scheduler.step(iou) if iou best_iou: best_iou iou torch.save(model.state_dict(), best_unet.pth) print(fepoch {epoch}, iou {iou:.4f}, best {best_iou:.4f})逻辑说明model.train()和model.eval()切换影响 BatchNorm 和 Dropout 行为分割任务里如果用了 BatchNorm验证时必须切 eval否则统计量会被验证集污染。scheduler.step(iou)放在验证之后modemax表示指标越大越好。保存state_dict()而不是整个模型加载时先实例化结构再load_state_dict这样换设备或改代码更灵活。断点续训就多存一个optimizer.state_dict()和 epoch 数。3.2 测试与推理单张图和批量预测的差别测试脚本通常做两件事加载权重对测试集或单张图输出预测掩码。单张图推理要手动加 batch 维度并且做和训练一致的预处理。批量预测则直接复用DataLoader但要注意shuffleFalse否则输出顺序和文件名对不上。下面是一个单张图推理并保存彩色掩码的写法import numpy as np import torch from PIL import Image def predict(image_path, model, device, size(256, 256)): model.eval() img Image.open(image_path).convert(RGB) img_resized img.resize(size) tensor F.to_tensor(img_resized).unsqueeze(0).to(device) with torch.no_grad(): logits model(tensor) pred torch.argmax(logits, dim1).squeeze(0).cpu().numpy() # 把类别索引映射成可视颜色 color_map np.array([[0, 0, 0], [255, 0, 0]], dtypenp.uint8) color_mask color_map[pred] Image.fromarray(color_mask).save(pred_color.png) return pred参数说明unsqueeze(0)增加 batch 维torch.no_grad()关闭梯度省显存argmax(dim1)在类别维取最大索引。color_map按你的类别数调整二分类两行多分类继续加。注意推理时的 resize 必须和训练一致训练用 256×256推理用原图直接送进去卷积核感受野对不上结果会明显变差。如果原图尺寸不固定常见做法是 pad 到 32 的倍数再推理最后裁回原尺寸。3.3 指标计算与可视化验证训练日志里的 IoU 只是数字真正判断模型有没有学会要把预测掩码和原图叠加看。代码包里如果有可视化脚本一般会输出原图、真实掩码、预测掩码三张并排图。IoU 的计算按类别做def compute_iou(pred, target, num_classes2): ious [] for cls in range(num_classes): pred_cls (pred cls) target_cls (target cls) intersection (pred_cls target_cls).sum() union (pred_cls | target_cls).sum() if union 0: continue ious.append(intersection / union) return sum(ious) / len(ious)逻辑说明union 0表示该类别在预测和真实里都没出现跳过避免除零。这个指标比像素准确率更能反映分割质量。验证时建议固定一组图每个 epoch 后都跑一遍肉眼对比变化比只看 loss 曲线直观得多。4. 避坑与排查这份代码包跑不起来时先查这几条4.1 现象loss 变成 NaN 或一直不降原因学习率太大、掩码像素值不连续、损失函数对越界类别没处理。解决先把学习率降到 1e-4打印掩码的唯一值确认是 0 到 num_classes-1如果掩码是 0/255用mask mask // 255或重新映射。Dice Loss 的smooth太小也可能在空掩码时炸调到 1e-5 到 1e-6 之间。4.2 现象验证 IoU 很高但预测图全黑或全白原因类别不平衡模型学会了全预测背景或全预测前景。解决在交叉熵里加weight参数给少数类更高权重或者把 Dice Loss 权重调大。另外检查argmax的维度是不是搞错了dim1是类别维写成dim0会得到完全错误的结果。4.3 现象训练时显存溢出原因batch_size太大、输入尺寸太大、没有用torch.no_grad()做验证。解决先把 batch 降到 2 或 1输入从 512 降到 256验证和推理阶段强制加with torch.no_grad():如果还不够用梯度累积模拟大 batch每几步才optimizer.step()一次。4.4 现象Windows 上num_workers大于 0 就卡死原因Windows 的 DataLoader 多进程和某些环境不兼容。解决把num_workers设为 0或者把训练代码放在if __name__ __main__:下面。Linux 上一般没这个问题但共享内存不足时也会报错可以减小batch_size或设置pin_memoryFalse。4.5 现象加载权重后预测结果和训练时不一致原因模型结构定义和保存时不一致或者预处理归一化参数不同。解决确认UNet的in_channels、num_classes、每层通道数完全一致检查推理时的均值和方差是不是和训练一样。常见做法是把预处理参数写进配置文件训练和推理都读同一份。5. 进阶技巧把这份 U-Net 改成你自己的分割任务拿到这份代码包后最值钱的用法不是原样跑一遍而是把它当成模板替换数据集和输出类别。第一步把你的标注掩码整理成单通道灰度图像素值从 0 开始连续编号背景为 0。第二步改SegDataset里的路径和类别数num_classes改成你的类别总数color_map同步扩展。第三步如果目标尺度差异大可以在编码器后加一个空洞空间金字塔池化模块或者把骨干换成预训练的 ResNet但要注意预训练权重输入是三通道输出层要重新初始化。第四步训练时先用小尺寸 128×128 快速验证流程能跑通再逐步加到 256 或 512这样能早发现数据对齐问题。第五步保存最优权重时同时存一份config.json记录输入尺寸、类别数、归一化参数下次推理直接读不用翻代码回忆。验证方法上我习惯在训练集里留两张图不参与训练每个 epoch 后跑一次预测把原图、真实掩码、预测掩码拼成一张图存下来。如果这两张图的预测从模糊逐渐变清晰、边界逐渐贴合说明模型真的在学如果一直没变化先回去查数据管道。另一个技巧是打印每个类别的 IoU而不是只看平均值平均值会掩盖小类别的崩溃。从那以后我每次拿到新的分割代码包都强制先跑一遍“单张图过拟合”测试拿一张图反复训练几十次看能不能把这张图完美分割出来。如果连单张图都过拟合不了说明网络结构或损失函数有根本问题不用浪费时间调参。希望帮到你。本文还有配套的精品资源点击获取