新闻详情

新闻详情

首页 / 资讯中心 / 详情

PyTorch实战:FCN与UNet语义分割从原理到部署

发布时间:2026/10/2 15:06:40来源:尧图网络
PyTorch实战:FCN与UNet语义分割从原理到部署
简介本资源面向具备一定深度学习基础的计算机视觉学习者与开发者聚焦PyTorch框架下UNet与FCN两种经典图像语义分割算法的完整实现与源码解析可用于课程设计、科研复现或工程入门。压缩包共18个文件约227KB以py脚本、ipynb交互式笔记为主辅以png预测可视化图、zbak备份文件及license、md说明文档覆盖模型定义、损失函数、训练与推理全流程。项目采用模块化设计将流程拆解为数据预处理、模型搭建、训练优化与预测验证四环节UNet以对称编码器-解码器配合跳跃连接实现多尺度特征融合FCN通过全卷积化支持任意尺寸图像端到端分割并引入混合精度训练、余弦退火与类别权重平衡等技巧。评估涵盖交并比、像素准确率等指标代码含类型注解与单元测试便于二次开发。目前已有81人学习下载适合希望系统掌握分割算法细节与工程规范的读者参考。1. 从一张广告牌说起UNet 与 FCN 到底能切出什么去年帮朋友处理一批户外广告牌巡检图需求很朴素把画面里的广告牌区域抠出来剩下的天空、树木、行人全部归为背景。我一开始想用目标检测框一下了事结果发现广告牌经常被灯杆、树枝遮挡矩形框里混进大量无关像素后续做尺寸测量和内容识别时误差大得离谱。这就是语义分割要解决的问题——不是画框而是给每个像素分类。PyTorch 生态里FCN 和 UNet 是两条最常被拿来落地的路线FCN 用全卷积替换全连接把分类网络改造成逐像素预测UNet 在此基础上加了编码器-解码器之间的跳跃连接让小目标和不规则边缘的定位精度明显提升。这份资源把两个算法的实现、训练脚本和推理入口都放在一起适合已经装好 PyTorch、想跑通第一个分割项目的人也适合需要对照源码理解跳跃连接、上采样和损失函数怎么配合的熟手。下面我按自己拆包复现的顺序把能抄的步骤和容易翻车的地方一次讲清。2. 环境与数据准备把 PyTorch 和分割数据集先对齐2.1 为什么分割项目对环境版本更敏感图像分割和普通分类任务不同它同时依赖卷积、上采样、转置卷积和逐像素损失这些算子在 CUDA 版本、cuDNN 版本和 PyTorch 小版本之间偶发行为差异。我遇到过同一份 UNet 代码在 PyTorch 1.12 上正常收敛换到 2.0 后转置卷积输出尺寸差一个像素导致标签和预测对不齐损失直接飙到 nan。常见做法是锁定一个经过验证的组合比如 Python 3.9 PyTorch 2.0.1 CUDA 11.8或者直接用 conda 创建独立环境避免和系统里已有的 pytorch 基础框架互相污染。如果你在 WSL 里搭环境注意把显卡驱动装在 Windows 侧WSL 内只装 CUDA toolkit 和 PyTorch否则会出现 torch.cuda.is_available() 返回 False 的玄学问题。# 创建独立环境避免和已有 pytorch 环境冲突 conda create -n unet_seg python3.9 -y conda activate unet_seg # 安装 PyTorch以 CUDA 11.8 为例具体命令按官网当前推荐调整 pip install torch2.0.1 torchvision0.15.2 --index-url https://download.pytorch.org/whl/cu118 # 验证 GPU 是否可用 python -c import torch; print(torch.__version__, torch.cuda.is_available())上面三条命令里第一条建环境第二条装框架第三条做自检。参数上唯一需要你改的是 CUDA 版本号它必须和本机驱动支持的最高版本匹配不是越高越好。如果输出是 False先别急着重装用 nvidia-smi 看驱动版本再对照 PyTorch 官网的兼容表多数情况是装成了 CPU 版。2.2 数据目录怎么组织才能直接喂给 Dataset这份源码默认读取的是 VOC 格式的分割数据集目录结构要求比较固定。我一般会先把自己的数据整理成下面这样再改 Dataset 里的路径而不是反过来去改代码逻辑。dataset/ ├── JPEGImages/ # 原图jpg 或 png ├── SegmentationClass/ # 标签图png像素值就是类别 id └── ImageSets/ └── Segmentation/ ├── train.txt # 每行一个文件名不带扩展名 └── val.txt标签图这里有个血泪经验很多人用 labelme 或 ps 导出标签时保存成了 RGB 三通道图每个类别颜色不同但像素值不是 0、1、2 这样的连续 id。UNet 训练时做交叉熵会直接拿像素值当类别索引结果就是类别对不上、损失不下降。正确做法是确保标签图是单通道 8 位背景为 0目标从 1 开始递增。如果手上只有彩色标签写个映射脚本转一下比在训练里硬编码颜色表可靠得多。import numpy as np from PIL import Image # 把 RGB 彩色标签转成单通道类别 id 图 color_map { (0, 0, 0): 0, # 背景 (255, 0, 0): 1, # 类别1 (0, 255, 0): 2, # 类别2 } def rgb_to_label(mask_path, save_path): rgb np.array(Image.open(mask_path).convert(RGB)) label np.zeros(rgb.shape[:2], dtypenp.uint8) for color, idx in color_map.items(): match np.all(rgb color, axis-1) label[match] idx Image.fromarray(label).save(save_path) rgb_to_label(dataset/SegmentationClass/0001.png, dataset/SegmentationClass/0001_label.png)这段脚本的核心是 color_map 字典键是标签图里实际出现的 RGB 值值是训练时使用的类别 id。跑之前先用取色工具确认几个像素点的颜色别凭肉眼猜。转换完抽查几张用 numpy.unique 看下像素值分布确认没有遗漏的颜色被归到背景。3. FCN 实现拆解全卷积、上采样与跳级融合怎么落地3.1 从分类网络到逐像素预测的关键改动FCN 的思路是把 VGG 或 ResNet 最后的全连接层换成卷积层让网络输出从一维向量变成二维特征图再通过上采样恢复到原图尺寸。源码里 FCN 的主干通常用预训练 VGG16前 13 层卷积保持不变后面接 1x1 卷积把通道数压到类别数最后用转置卷积放大 32 倍。这个 32 倍版本就是 FCN-32s边缘比较粗糙源码里还实现了 FCN-16s 和 FCN-8s分别把 pool4 和 pool3 的特征拿过来做融合小目标分割效果会好一些。选哪个版本取决于你的目标尺寸广告牌、道路这类大区域FCN-32s 够用细胞、裂缝这类细长目标直接上 FCN-8s别在 32s 上浪费时间调参。import torch import torch.nn as nn import torchvision.models as models class FCN8s(nn.Module): def __init__(self, num_classes): super().__init__() vgg models.vgg16(pretrainedTrue).features # 取 VGG 不同阶段的特征用于后续跳级融合 self.stage1 vgg[:17] # 到 pool3 self.stage2 vgg[17:24] # 到 pool4 self.stage3 vgg[24:31] # 到 pool5 self.score_pool4 nn.Conv2d(512, num_classes, 1) self.score_pool3 nn.Conv2d(256, num_classes, 1) self.upscore2 nn.ConvTranspose2d(num_classes, num_classes, 4, stride2, padding1) self.upscore8 nn.ConvTranspose2d(num_classes, num_classes, 16, stride8, padding4) def forward(self, x): p3 self.stage1(x) p4 self.stage2(p3) p5 self.stage3(p4) score self.score_pool4(p4) self.upscore2(self.score_pool3(p3)) return self.upscore8(score)这段代码里stage1、stage2、stage3 把 VGG 切成三段分别拿到 pool3、pool4、pool5 的特征。score_pool4 和 score_pool3 是两个 1x1 卷积把通道数统一到类别数方便相加。upscore2 做 2 倍上采样upscore8 做 8 倍上采样最终输出和输入同尺寸。参数上 num_classes 要和你标签里的类别数一致包含背景。转置卷积的 kernel_size、stride、padding 三个值必须配套改一个就要重算输出尺寸否则拼接时会报维度不匹配。3.2 训练循环里损失函数和指标怎么选分割任务最常用的损失是交叉熵但类别极不平衡时比如背景占 90%交叉熵会被背景主导模型学会全预测背景也能拿到高准确率。源码里默认用 CrossEntropyLoss我一般会改成带权重的版本或者叠加 Dice Loss。指标上别只看 pixel accuracy它在这个场景下会骗人重点看 mIoU 和每个类别的 IoU。import torch.nn.functional as F def train_one_epoch(model, loader, optimizer, device, weightNone): model.train() total_loss 0 for imgs, labels in loader: imgs, labels imgs.to(device), labels.to(device).long() optimizer.zero_grad() outputs model(imgs) # weight 用于给少数类别更高权重缓解不平衡 loss F.cross_entropy(outputs, labels, weightweight) loss.backward() optimizer.step() total_loss loss.item() return total_loss / len(loader)weight 是一个长度为 num_classes 的张量值越大表示该类越重要。计算方式常见的是取类别频率的倒数再归一化。如果你不想手算先用默认交叉熵跑几轮看验证集里小类别的 IoU 是不是接近 0是的话再加权重。学习率方面FCN 微调预训练主干时用 1e-4 比较稳从零训练可以到 1e-3但 batch size 要相应调整。4. UNet 实现拆解编码器-解码器与跳跃连接的真实作用4.1 下采样、上采样和拼接的尺寸对齐UNet 的结构像一个 U 形左边编码器不断卷积加池化特征图变小、通道变多右边解码器不断上采样加卷积特征图变大、通道变少中间用跳跃连接把编码器同层的特征直接拼到解码器对应层。这个拼接是 UNet 比 FCN 在医学图像分割、小目标分割上表现更好的核心原因因为它保留了下采样过程中丢失的空间细节。但拼接对尺寸要求很严格编码器第 n 层输出是 64x64解码器对应层上采样后也必须是 64x64差一个像素就报错。我一般会在拼接前打印两边 shape确认一致再往下走。class UNet(nn.Module): def __init__(self, in_channels3, num_classes2): super().__init__() def block(in_c, out_c): return nn.Sequential( nn.Conv2d(in_c, out_c, 3, padding1), nn.BatchNorm2d(out_c), nn.ReLU(inplaceTrue), nn.Conv2d(out_c, out_c, 3, padding1), nn.BatchNorm2d(out_c), nn.ReLU(inplaceTrue), ) self.enc1 block(in_channels, 64) self.enc2 block(64, 128) self.pool nn.MaxPool2d(2) self.bottleneck block(128, 256) self.up2 nn.ConvTranspose2d(256, 128, 2, stride2) self.dec2 block(256, 128) # 拼接后通道翻倍所以输入是 256 self.up1 nn.ConvTranspose2d(128, 64, 2, stride2) self.dec1 block(128, 64) self.out nn.Conv2d(64, num_classes, 1) def forward(self, x): e1 self.enc1(x) e2 self.enc2(self.pool(e1)) b self.bottleneck(self.pool(e2)) d2 self.up2(b) d2 self.dec2(torch.cat([e2, d2], dim1)) # 跳跃连接 d1 self.up1(d2) d1 self.dec1(torch.cat([e1, d1], dim1)) return self.out(d1)代码里 enc1、enc2 是编码器两个阶段bottleneck 是底部up2、up1 是转置卷积上采样dec2、dec1 是解码器卷积块。关键点在 dec2 的输入通道是 256因为 torch.cat 把 e2 的 128 通道和上采样后的 128 通道拼在一起了。如果你把 num_classes 改成 3 或更多只需要改 out 层其他不用动。BatchNorm 在小 batch size 下表现不稳定如果显存只够跑 batch size 2建议换成 GroupNorm 或 InstanceNorm。4.2 训练自己的数据集要改哪几个地方拿到源码后直接跑 demo 数据通常没问题但换成自己的数据至少要改三处Dataset 里的路径和类别数、模型初始化时的 num_classes、以及可视化时的颜色表。我见过有人只改了路径没改类别数训练不报错但预测全是背景排查半天才发现输出通道还是 2。另外UNet 对输入尺寸没有硬性要求但为了下采样和上采样能整除建议把图片 resize 到 16 的倍数比如 256x256 或 512x512。如果原图长宽比很重要用 padding 而不是直接拉伸。from torch.utils.data import Dataset, DataLoader from PIL import Image import torchvision.transforms as T class SegDataset(Dataset): def __init__(self, img_dir, mask_dir, file_list, size256): self.img_dir img_dir self.mask_dir mask_dir self.names [l.strip() for l in open(file_list)] self.size size self.img_tf T.Compose([T.Resize((size, size)), T.ToTensor()]) self.mask_tf T.Compose([T.Resize((size, size), interpolationT.InterpolationMode.NEAREST)]) def __len__(self): return len(self.names) def __getitem__(self, idx): name self.names[idx] img Image.open(f{self.img_dir}/{name}.jpg).convert(RGB) mask Image.open(f{self.mask_dir}/{name}.png) return self.img_tf(img), torch.from_numpy(np.array(self.mask_tf(mask))).long() loader DataLoader(SegDataset(dataset/JPEGImages, dataset/SegmentationClass, dataset/ImageSets/Segmentation/train.txt), batch_size4, shuffleTrue, num_workers2)标签图的 resize 必须用最近邻插值用双线性会把类别 id 插成小数再取整就乱了。num_workers 在 Windows 上有时会卡死设成 0 先跑通再逐步往上加。batch_size 根据显存调8G 显存跑 256x256 的 UNetbatch size 4 到 8 比较合适。5. 避坑与排查分割训练里最容易翻车的五件事5.1 损失不下降先查标签而不是模型现象训练几个 epochloss 在 0.69 附近震荡准确率不动。原因标签图里只有 0 和 255 两个值255 被当成类别 255而模型输出只有 2 类交叉熵计算时目标越界实际梯度是乱的。解决用 numpy.unique 检查标签像素值确保是 0 到 num_classes-1 的连续整数255 要映射成 1。5.2 验证集 IoU 很高但预测图全黑现象mIoU 显示 0.9但把预测结果叠到原图上看目标区域根本没被标出来。原因背景占比太高模型全预测背景背景 IoU 拉高了平均值小类别 IoU 是 0。解决分开打印每个类别的 IoU别只看平均值同时把验证集里目标占比高的样本单独抽出来看。5.3 转置卷积输出尺寸和输入对不上现象UNet 拼接时 RuntimeError: Sizes of tensors must match。原因输入图片尺寸不是 16 的倍数经过多次下采样后奇数尺寸被取整上采样回来差一个像素。解决在 Dataset 里统一 resize 到 16 的倍数或者在拼接前用 F.interpolate 把上采样结果对齐到编码器特征的尺寸。5.4 显存够但训练速度极慢现象GPU 利用率只有 20%一个 epoch 要跑十几分钟。原因num_workers 设成 0数据加载在主进程串行执行GPU 一直在等数据。解决把 num_workers 调到 CPU 核心数的一半左右同时开 pin_memoryTrue让数据搬运和计算重叠。5.5 推理结果比训练时差很多现象训练时指标正常用单张图推理时边缘破碎、类别错乱。原因推理时忘了加 model.eval()BatchNorm 还在用当前 batch 的统计量或者输入没有做和训练一致的归一化。解决推理前固定写 model.eval() 和 torch.no_grad()并把训练时的均值和标准差抄过来做 Normalize。6. 进阶技巧把 UNet 推理结果导出 ONNX 并验证一致性训练完模型只是第一步真正部署时经常要把 PyTorch 模型转成 ONNX再交给推理引擎或板端。这一步最容易出的问题是导出成功但结果对不上所以我会强制做一次数值比对。下面这段脚本把 UNet 导出成 ONNX然后用 onnxruntime 跑同一张输入和 PyTorch 输出比最大绝对误差。import torch import numpy as np import onnxruntime as ort model UNet(num_classes2) model.load_state_dict(torch.load(best_unet.pth, map_locationcpu)) model.eval() dummy torch.randn(1, 3, 256, 256) torch.onnx.export(model, dummy, unet.onnx, input_names[input], output_names[output], opset_version11, dynamic_axes{input: {0: batch}}) with torch.no_grad(): torch_out model(dummy).numpy() sess ort.InferenceSession(unet.onnx) onnx_out sess.run(None, {input: dummy.numpy()})[0] print(max abs diff:, np.abs(torch_out - onnx_out).max())导出时 opset_version 建议用 11 或更高低版本对某些上采样算子支持不好。dynamic_axes 把 batch 维设成动态部署时可以传不同 batch size。比对结果里 max abs diff 在 1e-4 以内算正常超过 1e-2 就要查是不是有算子被替换成了近似实现。我一般还会把 ONNX 模型的输出 argmax 成类别图和 PyTorch 的 argmax 结果逐像素比确认没有大面积翻转。这套流程走完模型才算真正能交出去。从那以后我每次导出 ONNX 都强制跑一遍数值比对不再凭“导出没报错”就认为没问题。希望帮到你。本文还有配套的精品资源点击获取
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

CTU-13数据集深度解析:网络入侵检测的黄金标定基准 2026/10/2 15:51:39

CTU-13数据集深度解析:网络入侵检测的黄金标定基准

1. 项目概述:CTU-13数据集不是“新发布的网红数据集”,而是网络流量分析领域里一块被反复打磨十年的“校准砝码”你如果最近在查“CTU-13数据集”,大概率是刚接触网络入侵检测、流量行为建模或AI安全方向,被某篇论文、某门课程作业…

阅读更多 →
手机App开发方案落地:从技术选型到MVP构建的完整指南 2026/10/2 15:51:39

手机App开发方案落地:从技术选型到MVP构建的完整指南

简介:这是一份面向房地产企业营销团队、产品经理及移动应用开发者的APP开发方案借鉴资料,聚焦如何用手机App重构传统楼书与购房沟通方式。方案提出随身楼书、多媒体展示、信息实时推送等核心思路,并系统拆解出楼盘介绍、周边配套、房型展示、…

阅读更多 →
OpenRig开源项目:用铝型材DIY高刚性模拟驾驶舱全指南 2026/10/2 15:51:39

OpenRig开源项目:用铝型材DIY高刚性模拟驾驶舱全指南

OpenRig 是今年我在模拟赛车圈里正式发起的一个开源项目,目标很直接:用一套标准欧标铝型材和通用连接件,搭出一台结构刚性足够、尺寸可调、成本只有成品三分之一左右的模拟驾驶舱(sim rig)。玩模拟赛车的朋友应该都懂&…

阅读更多 →
英语教学AI引擎:可干预、可回溯、可评估的情景教学系统 2026/10/2 15:51:39

英语教学AI引擎:可干预、可回溯、可评估的情景教学系统

1. 这不是又一个“AI聊天框”,而是一套能进课堂的英语教学引擎 我第一次在中学试讲时,把刚写好的英语情景对话Agent投到投影仪上,学生没点开就笑了:“老师,这回是不是又要听机器人念课文?”——结果三分钟后…

阅读更多 →
RAG从原理到实战:用客服项目讲透检索增强生成全链路 2026/10/2 15:51:39

RAG从原理到实战:用客服项目讲透检索增强生成全链路

上周有位读者找我,说面试官让他用一分钟解释RAG,他背了几天八股还是卡壳。这事我特别理解:RAG概念本身不难,难的是你手里没有一条完整的实现链路,脑子里只有“向量检索 大模型”六个字,自然说不清楚。今天…

阅读更多 →
QMenu 删除崩溃现象及解决方法:从 delete 到信号槽的完整排查 2026/10/2 15:51:32

QMenu 删除崩溃现象及解决方法:从 delete 到信号槽的完整排查

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞 ✉