news 2026/9/26 4:15:46

PyTorch实战:SegNet图像分割源码解析与训练避坑指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch实战:SegNet图像分割源码解析与训练避坑指南

简介:这份资源是基于PyTorch实现SegNet图像分割任务的完整Python源码包,面向计算机相关专业正在做课程设计、期末大作业的学生,以及需要项目实战练习的学习者。项目经导师指导并认可,获得98分成绩,可作为图像语义分割方向的参考方案。压缩包共119个文件,约27.19MB,包含14个py源码文件、77个png图像数据、12个pyc编译文件,以及pth模型权重、sh运行脚本、Dockerfile、env环境配置、logging.ini日志配置、README.md说明文档和pdf资料等,覆盖从数据、模型到部署的完整链路。目前已有175人学习下载。读者可从中获取SegNet网络结构搭建、编码器-解码器实现、训练与推理流程、日志记录及容器化运行等关键代码,并借助模型权重与配置脚本快速复现实验,理解图像分割任务的工程组织方式与排错思路,适合作为大作业模板或进阶练手项目。

1. 一份能直接跑通的 SegNet 图像分割源码,到底解决了什么

如果你正在为期末大作业或者课程设计找一份能跑、能改、能写进报告的图像分割代码,那这份基于 PyTorch 实现 SegNet 的 Python 源码包,大概率能省掉你从零搭网络结构的两三天时间。SegNet 本身是图像分割里一个非常经典的编码器-解码器结构,核心卖点是用最大池化索引来做上采样,相比全卷积网络那种反卷积方式,参数量更小、边界恢复更准,在道路分割、医学图像分割这类像素级分类任务里一直是教学和工程入门的首选。这份源码把数据加载、模型定义、训练循环、指标计算和推理可视化都串起来了,适合刚接触 PyTorch 图像分割的本科生,也适合想快速验证自己数据集的从业者。你拿到手之后,改数据路径、调几个超参就能跑,不用再纠结编码器解码器怎么对齐、池化索引怎么传这些细节。

2. SegNet 的编码器-解码器结构:为什么池化索引是它的命门

2.1 从 VGG 骨干到对称解码:结构拆解

SegNet 的整体结构可以理解成一条“下坡再上坡”的路。编码器部分直接沿用 VGG16 的前 13 层卷积,分成 5 个 stage,每个 stage 里堆两到三个 3×3 卷积加 ReLU,然后接一个 2×2 最大池化。每次池化,特征图长宽减半,通道数翻倍,最终把一张 H×W×3 的图压成 H/32×W/32×512 的特征块。解码器则是完全对称的 5 个 stage,每一步先做上采样把长宽翻倍,再堆卷积把通道数降回去,最后接一个 1×1 卷积输出类别数通道,softmax 之后就是每个像素的类别概率。

关键差异在于上采样方式。很多分割网络用反卷积或者双线性插值,SegNet 用的是“最大池化索引”。编码器每次做 2×2 最大池化时,不光输出最大值,还记录下最大值在 2×2 窗口里的位置(0 到 3 的索引)。解码器上采样时,直接把值填回对应位置,其余位置补零。这样做的好处是边界信息不会在插值里被抹平,而且不需要学习上采样参数,显存占用比反卷积小一截。常见做法是在编码器里用一个MaxPool2d的return_indices=True,把索引一路存下来传给解码器。

2.2 用 PyTorch 把编码器和解码器搭出来

下面这段代码是 SegNet 编码器一个 stage 的典型写法,我一般会把它封装成EncoderBlock,方便复用。

import torch import torch.nn as nn class EncoderBlock(nn.Module): def __init__(self, in_ch, out_ch, num_convs=2): super().__init__() layers = [] for i in range(num_convs): layers.append(nn.Conv2d(in_ch if i == 0 else out_ch, out_ch, kernel_size=3, padding=1)) layers.append(nn.BatchNorm2d(out_ch)) layers.append(nn.ReLU(inplace=True)) self.conv = nn.Sequential(*layers) # return_indices=True 是 SegNet 的灵魂,必须开 self.pool = nn.MaxPool2d(kernel_size=2, stride=2, return_indices=True) def forward(self, x): x = self.conv(x) x, indices = self.pool(x) return x, indices

逻辑说明:num_convs控制这个 stage 里堆几个卷积,VGG16 的 5 个 stage 分别是 2、2、3、3、3。BatchNorm2d加在卷积和 ReLU 之间,能明显稳住训练初期的 loss 震荡。return_indices=True让池化层多返回一个索引张量,形状和输出特征图一致,后面解码器要用它做MaxUnpool2d。

参数说明:in_ch是输入通道,第一个 stage 是 3,后面依次是 64、128、256、512。out_ch对应 64、128、256、512、512。kernel_size=3, padding=1保证卷积不改变长宽,只有池化在降分辨率。

解码器这边对应写一个DecoderBlock,核心是MaxUnpool2d接收编码器传来的索引。

class DecoderBlock(nn.Module): def __init__(self, in_ch, out_ch, num_convs=2): super().__init__() self.unpool = nn.MaxPool2d(kernel_size=2, stride=2, return_indices=True) # 占位,实际用 MaxUnpool2d self.unpool = nn.MaxUnpool2d(kernel_size=2, stride=2) layers = [] for i in range(num_convs): layers.append(nn.Conv2d(in_ch if i == 0 else out_ch, out_ch, kernel_size=3, padding=1)) layers.append(nn.BatchNorm2d(out_ch)) layers.append(nn.ReLU(inplace=True)) self.conv = nn.Sequential(*layers) def forward(self, x, indices, output_size): x = self.unpool(x, indices, output_size=output_size) x = self.conv(x) return x

逻辑说明:MaxUnpool2d的前向需要三个东西——当前特征图、编码器对应层的索引、以及上采样后的目标尺寸。目标尺寸在 PyTorch 里可以用output_size显式指定,避免因为奇数尺寸导致形状对不上。output_size一般直接传编码器池化前的特征图尺寸,可以在编码器 forward 里顺手存下来。

参数说明:in_ch和out_ch跟编码器反着来,解码器第一个 stage 是 512 进 512 出,最后是 64 进 64 出。num_convs同样对应 3、3、3、2、2。整个解码器最后接一个Conv2d(64, num_classes, kernel_size=1)输出类别 logits。

2.3 完整前向流程与输出尺寸对齐

把编码器和解码器串起来的时候,最容易翻车的地方就是尺寸对不上。假设输入是 360×480,经过 5 次池化变成 12×15,解码器第一次上采样要回到 24×30,第二次 48×60,第三次 96×120,第四次 192×240,第五次 384×480。如果输入尺寸不能被 32 整除,最后一次上采样出来的尺寸就会和原图差几个像素,算 loss 的时候直接报形状错误。

我一般会在数据预处理里强制 resize 到 32 的倍数,比如 352×480 或者 384×512。如果不想改数据,也可以在解码器最后加一个F.interpolate把输出拉回原图尺寸,但这样会引入额外的插值误差,边界精度会掉一点。常见做法是训练时 resize 到固定尺寸,推理时再插值回原图,这样训练稳定、推理灵活。

3. 数据管道与训练循环:从文件夹到可收敛的模型

3.1 数据集组织与 Dataset 类写法

图像分割的数据集通常有两种组织方式:一种是原图和掩码图分两个文件夹,文件名一一对应;另一种是原图和掩码图放在同一个文件夹,用后缀区分。这份源码一般会采用第一种,目录结构像这样:

dataset/ images/ 0001.png 0002.png masks/ 0001.png 0002.png

掩码图是单通道的 PNG,每个像素值就是类别 id,背景是 0,目标类别从 1 开始。写Dataset类的时候,关键是把原图和掩码用相同的文件名读进来,然后做同步的随机增强。同步增强是分割任务里最容易忽略的坑,图像翻转了掩码没翻,训练出来的模型直接学废。

import os from PIL import Image import torch from torch.utils.data import Dataset import torchvision.transforms.functional as TF import random class SegDataset(Dataset): def __init__(self, root, split='train', size=(352, 480)): self.img_dir = os.path.join(root, 'images') self.mask_dir = os.path.join(root, 'masks') self.names = sorted(os.listdir(self.img_dir)) self.size = size 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') # 同步 resize img = TF.resize(img, self.size) mask = TF.resize(mask, self.size, interpolation=Image.NEAREST) # 同步随机水平翻转 if random.random() > 0.5: img = TF.hflip(img) mask = TF.hflip(mask) img = TF.to_tensor(img) img = TF.normalize(img, mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) mask = torch.from_numpy( __import__('numpy').array(mask, dtype='int64')) return img, mask

逻辑说明:convert('L')把掩码转成单通道,resize时掩码必须用NEAREST插值,否则类别 id 会被插值成小数,后面算交叉熵直接报错。水平翻转对图像和掩码同时做,保证空间对应关系不变。归一化用 ImageNet 的均值和标准差,因为编码器是 VGG 预训练权重,输入分布最好对齐。

参数说明:size建议设成 32 的倍数,(352, 480)是一个比较稳的选择,显存占用和精度平衡得不错。如果显存只有 6GB,可以降到(256, 352)。num_classes根据你的数据集来,二分类就是 2,多分类就改成实际类别数。

3.2 损失函数与优化器配置

分割任务最常用的损失是交叉熵,PyTorch 里用nn.CrossEntropyLoss,它内部会做 softmax,所以模型输出直接给 logits 就行,不要在模型里加 softmax。如果类别极度不平衡,比如背景占 90% 以上,可以给weight参数传一个按类别频率倒数算出来的权重张量,或者换成 Dice Loss、Focal Loss。我一般会先用交叉熵跑一版 baseline,看每个类别的 IoU,再决定要不要换损失。

优化器用 Adam 或者 SGD 都行。Adam 收敛快,适合快速验证;SGD 加 momentum 泛化稍好,适合最终刷指标。学习率初始设 1e-3(Adam)或者 1e-2(SGD),配合StepLR或者CosineAnnealingLR衰减。batch size 在 8 到 16 之间,取决于显存。

import torch.optim as optim from torch.optim.lr_scheduler import CosineAnnealingLR model = SegNet(num_classes=2).cuda() criterion = nn.CrossEntropyLoss(ignore_index=255) # 255 是忽略像素 optimizer = optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4) scheduler = CosineAnnealingLR(optimizer, T_max=50)

逻辑说明:ignore_index=255用来忽略掩码里标记为“不确定”或者“边界”的像素,这些像素不参与 loss 计算,能避免模型被噪声标签带偏。weight_decay加一点 L2 正则,防止过拟合。CosineAnnealingLR的T_max设成总 epoch 数,学习率会从 1e-3 平滑降到接近 0。

参数说明:ignore_index要和你的掩码标注约定一致,如果掩码里没有 255 这个值,可以去掉这个参数。weight_decay一般设 1e-4 到 1e-5,太大模型欠拟合,太小正则效果不明显。

3.3 训练循环与验证指标

训练循环的骨架很固定:前向、算 loss、反向、更新。每个 epoch 结束后在验证集上算 mIoU 和像素准确率。mIoU 是分割任务最核心的指标,计算方式是每个类别的交并比取平均。

def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss = 0 for imgs, masks in loader: imgs, masks = imgs.to(device), masks.to(device) optimizer.zero_grad() outputs = model(imgs) loss = criterion(outputs, masks) loss.backward() optimizer.step() total_loss += loss.item() return total_loss / len(loader) @torch.no_grad() def evaluate(model, loader, num_classes, device): model.eval() hist = torch.zeros(num_classes, num_classes) for imgs, masks in loader: imgs, masks = imgs.to(device), masks.to(device) outputs = model(imgs) preds = outputs.argmax(dim=1) # 混淆矩阵统计 for t, p in zip(masks.view(-1), preds.view(-1)): hist[t.long(), p.long()] += 1 iou = torch.diag(hist) / (hist.sum(1) + hist.sum(0) - torch.diag(hist)) return iou.nanmean().item()

逻辑说明:argmax(dim=1)把每个像素的类别概率转成类别 id。混淆矩阵hist的行是真实类别,列是预测类别,对角线是预测正确的像素数。IoU 用diag / (行和 + 列和 - 对角)算,nanmean忽略没有出现的类别。

参数说明:num_classes要和模型输出通道一致。如果某个类别在验证集里一个像素都没有,它的 IoU 会是 nan,nanmean会自动跳过。验证时记得model.eval()和torch.no_grad(),省显存也省时间。

4. 推理、可视化与模型导出:把结果变成能交差的图

4.1 单张图像推理与掩码上色

训练完之后,最直观的验证方式就是拿几张测试图跑一遍,把预测掩码上色后和原图并排显示。上色可以用一个固定的颜色表,每个类别对应一个 RGB 值。

import numpy as np import matplotlib.pyplot as plt def colorize_mask(mask, palette): h, w = mask.shape color = np.zeros((h, w, 3), dtype=np.uint8) for cls_id, rgb in enumerate(palette): color[mask == cls_id] = rgb return color palette = [(0, 0, 0), (255, 0, 0), (0, 255, 0), (0, 0, 255)] model.eval() img, _ = dataset[0] with torch.no_grad(): pred = model(img.unsqueeze(0).cuda()).argmax(1).squeeze().cpu().numpy() colored = colorize_mask(pred, palette) plt.imshow(colored) plt.savefig('pred.png')

逻辑说明:palette的长度等于类别数,索引就是类别 id。colorize_mask用布尔索引批量上色,比逐像素循环快很多。推理时unsqueeze(0)加一个 batch 维度,因为模型 forward 期望 4D 输入。

参数说明:palette的颜色可以自己定,但建议背景用黑色,目标类别用高对比度颜色,方便肉眼检查。如果类别多,可以用matplotlib的tab20色表生成。

4.2 导出 ONNX 与 TorchScript

如果作业要求部署或者跨平台推理,可以把模型导出成 ONNX 或者 TorchScript。ONNX 的好处是可以用 ONNX Runtime 在 CPU 上跑,不依赖 PyTorch 环境。

dummy = torch.randn(1, 3, 352, 480).cuda() torch.onnx.export(model, dummy, 'segnet.onnx', input_names=['input'], output_names=['output'], opset_version=11, dynamic_axes={'input': {0: 'batch'}})

逻辑说明:dummy是一个示例输入,用来追踪计算图。opset_version=11兼容性比较好,dynamic_axes把 batch 维度设成动态,导出后的模型可以接受任意 batch size。

参数说明:如果模型里有MaxUnpool2d,ONNX 对它的支持在 opset 11 之后才完善,所以不要用太低的版本。导出前确保模型在 eval 模式,否则 BatchNorm 的统计量会不对。

5. 避坑与排查:SegNet 训练里最容易翻车的五个地方

5.1 现象:loss 一直是 nan,训练几个 step 就崩

原因:学习率太大,或者输入没有归一化,导致梯度爆炸。SegNet 编码器是 VGG 结构,对输入尺度很敏感,如果图像像素值还是 0 到 255,第一层卷积输出会非常大。

解决:确认ToTensor()把像素值压到 0 到 1,再做 normalize。学习率从 1e-4 开始试,如果还 nan,加梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)。

5.2 现象:mIoU 一直卡在 0.2 左右上不去

原因:掩码的类别 id 没对齐,或者ignore_index设错了。比如掩码里背景是 0、目标是 255,但ignore_index设成了 0,结果背景全被忽略,模型只学目标,IoU 自然低。

解决:用numpy.unique打印掩码里所有出现的像素值,确认类别 id 范围。ignore_index只设成真正需要忽略的值,不要误伤背景。

5.3 现象:解码器上采样后尺寸和编码器对不上,报形状错误

原因:输入图像尺寸不能被 32 整除,五次池化后出现奇数尺寸,MaxUnpool2d恢复出来的尺寸和编码器池化前差一个像素。

解决:在 Dataset 里强制 resize 到 32 的倍数,比如 352×480、384×512。如果必须保持原尺寸,在解码器最后加F.interpolate拉回原图,但要注意这会影响边界精度。

5.4 现象:验证集 loss 比训练集低很多,但 mIoU 也很低

原因:验证集的掩码里有很多ignore_index像素,导致 loss 被低估,但实际预测的像素很少,IoU 自然低。这种情况常见于边界标注很粗的数据集。

解决:检查验证集掩码的ignore_index比例,如果超过 20%,说明标注质量有问题,要么重新标注,要么在 loss 里给有效像素加权。

5.5 现象:显存不够,batch size 降到 2 还是 OOM

原因:SegNet 编码器最大通道 512,解码器对称,中间特征图在 352×480 输入下占用不小。如果还开了return_indices,索引张量也要占显存。

解决:把输入尺寸降到 256×352,或者把编码器前两个 stage 的通道数减半。也可以用混合精度训练torch.cuda.amp,显存能省 30% 到 40%,速度还快一点。

from torch.cuda.amp import autocast, GradScaler scaler = GradScaler() with autocast(): outputs = model(imgs) loss = criterion(outputs, masks) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()

逻辑说明:autocast自动把部分运算转成 float16,GradScaler防止梯度下溢。这套组合在 6GB 显存的卡上能把 batch size 从 2 提到 6 左右。

参数说明:混合精度训练对MaxUnpool2d的支持没问题,但CrossEntropyLoss在 autocast 下会自动转回 float32,不用手动改。

6. 把 SegNet 用到自己的数据上:迁移学习与类别不平衡处理

拿到这份源码之后,最实际的用法不是从头训练,而是加载 VGG16 预训练权重做迁移学习。SegNet 编码器结构和 VGG16 前 13 层完全一致,可以直接把torchvision.models.vgg16(pretrained=True)的 features 部分权重拷过来。这样即使你的数据集只有几百张图,也能在 30 个 epoch 内收敛到一个可用的 IoU。

import torchvision.models as models vgg = models.vgg16(pretrained=True) encoder_dict = model.encoder.state_dict() vgg_features = vgg.features.state_dict() # 只拷贝编码器里卷积层的权重 pretrained_dict = {k: v for k, v in vgg_features.items() if k in encoder_dict and 'weight' in k or 'bias' in k} encoder_dict.update(pretrained_dict) model.encoder.load_state_dict(encoder_dict)

逻辑说明:vgg.features里包含卷积层和池化层,但 SegNet 的编码器把池化单独拆出来了,所以只拷贝卷积和 BN 的权重。strict=False可以避免键名不匹配报错,但这里手动过滤更稳。

参数说明:pretrained=True会下载 ImageNet 权重,第一次运行需要联网。如果环境不能联网,可以提前把权重文件下好放到~/.cache/torch/hub/checkpoints/。

类别不平衡是分割任务里另一个绕不开的问题。医学图像分割里病灶往往只占几个像素,道路分割里车道线也很细。除了给CrossEntropyLoss加权重,还可以用 Dice Loss 或者 Tversky Loss。Dice Loss 直接优化预测和真实掩码的重叠度,对小目标更友好。

class DiceLoss(nn.Module): def __init__(self, smooth=1.0): super().__init__() self.smooth = smooth def forward(self, logits, targets): probs = torch.softmax(logits, dim=1) targets_onehot = torch.nn.functional.one_hot( targets, num_classes=probs.shape[1]).permute(0, 3, 1, 2).float() intersection = (probs * targets_onehot).sum(dim=(0, 2, 3)) union = probs.sum(dim=(0, 2, 3)) + targets_onehot.sum(dim=(0, 2, 3)) dice = (2 * intersection + self.smooth) / (union + self.smooth) return 1 - dice.mean()

逻辑说明:one_hot把掩码转成和 probs 同形状的 one-hot 张量,permute把类别维挪到通道维。intersection和union按类别求和,smooth防止除零。最终返回1 - dice作为 loss。

参数说明:smooth一般设 1.0,太小对空类别没效果,太大 loss 会被平滑掉。Dice Loss 可以和交叉熵按 0.5:0.5 加权组合,收敛更稳。

我自己的习惯是,每次拿到一份新的分割源码,先不急着改模型,而是用一张图跑一遍前向,把每一层的输出尺寸打印出来,确认编码器和解码器能对上。然后再用 10 张图过拟合一遍,如果 loss 能降到接近 0,说明模型和损失函数没问题,剩下的就是调数据和超参。这套流程帮我省掉了很多次盲目调参的时间。从那以后我每次跑新数据集都强制走一遍“单图前向 + 小样本过拟合”,希望帮到你。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/26 4:15:29

尖音符´:从Unicode到编程避坑的完整指南

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

作者头像 李华
网站建设 2026/9/26 4:14:45

[Git-1] Git基础认识

一、版本控制 Git 是软件开发中最常用的版本控制工具之一。可能大家都有接触过,不过也就仅限于常用命令的使用,其原理并不是很清楚。 1、版本控制引入 在没有版本控制工具时,如果我们想保存代码的不同版本,最直接的方法可能就是复…

作者头像 李华
网站建设 2026/9/26 4:14:23

星阅书城接口自动化测试实战揭秘

1. 测试概述1.1 测试背景星阅书城平台的登录、用户管理、商品与订单等接口是业务主链路:用户必须先登录拿到Token凭证,才能新增、修改用户,下单链路则要按 "商品列表 → 商品详情 → 提交订单 → 订单支付 → 校验订单状态" 的顺序…

作者头像 李华
网站建设 2026/9/26 4:14:05

终端的庖丁解牛

根因 早期计算机是大型主机,主机本体放在机房,运算、存储全部由主机完成。操作人员不可能趴在主机上操作,于是就造出终端:本身没有算力,只负责接收人的输入、展示主机返回的输出,通过线缆连接远端主机。 到…

作者头像 李华