ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

PyTorch实战:DeepLabV3在Cityscapes上的语义分割训练与避坑指南

2026/9/26 14:06:35 拓冰建站 浏览量
PyTorch实战:DeepLabV3在Cityscapes上的语义分割训练与避坑指南 简介这份资源面向计算机视觉方向的研究者、算法工程师与深度学习学习者提供在Cityscapes数据集上训练DeepLabV3语义分割模型的完整PyTorch实现帮助读者理解ASPP空洞空间金字塔池化与全局上下文模块的设计思路并掌握从数据预处理到模型评估的全流程。压缩包共18个文件以12个Python脚本为核心涵盖模型定义、训练、评估与数据加载等模块另含4个pth预训练权重、1个md说明文档与1个license许可文件整体约258.23MB目录结构清晰便于按模块查阅与二次开发。目前已有2094人学习下载适合希望快速复现基线、改进分割精度或深入理解DeepLabV3内部机制的读者参考。1. 语义分割落地为什么绕不开 DeepLabV3 Cityscapes如果你手头有一批街景、园区或道路图像需要把每个像素分到「车、人、路面、建筑」这些类别里那语义分割就是绕不开的一环。而 Cityscapes 是这套任务里最常被拿来当基准的数据集30 个类别、5000 张精细标注图覆盖城市道路场景。DeepLabV3 则是把空洞卷积和 ASPP 模块做到工程上足够稳的经典结构PyTorch 实现版本多、改起来方便适合做二次开发和迁移。这份资源就是一套在 Cityscapes 上训练 DeepLabV3 的 PyTorch 代码包含数据加载、模型定义、训练循环和推理脚本。它解决的不是「从零教你 PyTorch」的问题而是让你跳过环境折腾和结构拼装直接跑通一条完整的训练链路。适合已经装好 PyTorch、想快速验证分割效果的人也适合拿它当骨架改自己数据集的从业者。下面按「资源是什么 → 怎么用 → 坑在哪」的顺序拆开讲。2. 环境搭建与数据准备从 pytorch 安装到 Cityscapes 目录结构2.1 环境依赖与 pytorch 安装的版本选择这套代码对 PyTorch 版本不算挑剔但 CUDA 版本和显卡驱动要对上。常见做法是用 conda 建一个独立环境避免和系统里的 python 安装冲突。如果你用的是较新的显卡比如 7900xtx 这类在 WSL 下跑 pytorch 环境搭建也能走通只是要确认 ROCm 或 CUDA 的适配情况。# 创建独立环境python 版本建议 3.8 到 3.10 conda create -n deeplab python3.9 -y conda activate deeplab # 安装 pytorch以 CUDA 11.8 为例具体命令去 pytorch 官网按你的驱动选 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 # 其余依赖 pip install numpy opencv-python pillow tqdm tensorboard这里的关键参数是--index-url它决定你下载的是 CPU 版还是 GPU 版。很多人 pytorch 安装教程超详细看了一堆结果装完torch.cuda.is_available()返回 False多半是驱动版本和 CUDA 版本没对齐。装完先跑一句验证import torch print(torch.__version__) print(torch.cuda.is_available()) # 期望 True print(torch.cuda.get_device_name(0))如果第二行是 False别急着改代码先查驱动。这是最常见的翻车点后面避坑章节会细说。2.2 Cityscapes 数据集的下载与目录组织Cityscapes 官方提供两个包leftImg8bit是原图gtFine是标注。训练用 train 划分验证用 val 划分。下载后解压目录结构要整理成代码能识别的形式cityscapes/ ├── leftImg8bit/ │ ├── train/ │ │ ├── aachen/ │ │ │ ├── aachen_000000_000019_leftImg8bit.png │ │ └── ... │ └── val/ └── gtFine/ ├── train/ │ ├── aachen/ │ │ ├── aachen_000000_000019_gtFine_labelIds.png │ │ ├── aachen_000000_000019_gtFine_instanceIds.png │ │ └── ... └── val/代码里读的是labelIds.png不是labelTrainIds.png这两个容易搞混。labelIds是原始类别 IDlabelTrainIds是映射到 19 个训练类别的版本。如果你直接用labelTrainIds需要改数据加载里的映射逻辑否则类别对不上训练 loss 会异常。2.3 数据加载器的参数配置数据加载部分通常继承torch.utils.data.Dataset核心是__getitem__里做同步的随机裁剪和归一化。下面是一个典型的加载逻辑import os import torch import numpy as np from PIL import Image from torch.utils.data import Dataset, DataLoader import torchvision.transforms as T class CityscapesDataset(Dataset): def __init__(self, root, splittrain, crop_size(512, 1024)): self.root root self.split split self.crop_size crop_size self.images [] self.labels [] # 遍历目录收集文件对 img_dir os.path.join(root, leftImg8bit, split) lbl_dir os.path.join(root, gtFine, split) for city in os.listdir(img_dir): for f in os.listdir(os.path.join(img_dir, city)): if f.endswith(_leftImg8bit.png): img_path os.path.join(img_dir, city, f) lbl_name f.replace(_leftImg8bit.png, _gtFine_labelIds.png) lbl_path os.path.join(lbl_dir, city, lbl_name) if os.path.exists(lbl_path): self.images.append(img_path) self.labels.append(lbl_path) def __len__(self): return len(self.images) def __getitem__(self, idx): img Image.open(self.images[idx]).convert(RGB) lbl Image.open(self.labels[idx]) # 同步随机裁剪保证图像和标签对齐 i, j, h, w T.RandomCrop.get_params(img, self.crop_size) img T.functional.crop(img, i, j, h, w) lbl T.functional.crop(lbl, i, j, h, w) img T.functional.to_tensor(img) img T.functional.normalize(img, mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) lbl torch.from_numpy(np.array(lbl)).long() return img, lblcrop_size设成 (512, 1024) 是显存和精度的折中。Cityscapes 原图是 1024x2048直接整图训练对显存要求高裁剪后 batch size 能开到 8 到 16。normalize用的 ImageNet 均值方差因为 DeepLabV3 主干通常是在 ImageNet 上预训练的。标签转long是因为交叉熵损失要求 int64。如果这里忘了同步裁剪图像和标签会错位训练出来的模型预测全是乱的这是血泪经验。3. DeepLabV3 模型结构与 ASPP 模块的实现细节3.1 主干网络选型ResNet 还是 MobileNetDeepLabV3 本身是个元结构主干可以换。代码里常见的是 ResNet-101 和 MobileNetV2 两种。ResNet-101 精度高但显存吃紧MobileNetV2 轻量适合边缘部署。选哪个取决于你的场景如果只是验证算法效果用 ResNet-101如果要往移动端推MobileNetV2 更实际。import torchvision.models as models import torch.nn as nn def build_backbone(nameresnet101, pretrainedTrue): if name resnet101: model models.resnet101(pretrainedpretrained) # 去掉最后的全连接和平均池化保留卷积特征 return nn.Sequential(*list(model.children())[:-2]), 2048 elif name mobilenetv2: model models.mobilenet_v2(pretrainedpretrained) return model.features, 1280 else: raise ValueError(fUnsupported backbone: {name})pretrainedTrue会下载 ImageNet 预训练权重第一次跑需要联网。返回的通道数 2048 或 1280 要传给 ASPP 模块作为输入通道。如果这里通道数写错后面 ASPP 的卷积层会直接报维度不匹配。3.2 ASPP 模块的空洞卷积率设置ASPP 是 DeepLabV3 的核心用不同空洞率的卷积并行提取多尺度特征。标准配置是 rates[6, 12, 18]加上一个全局平均池化分支。空洞率的选择和输入尺寸有关如果裁剪尺寸改成 256x512rates 也要相应调小否则感受野超出图像边界padding 会引入大量无效信息。class ASPP(nn.Module): def __init__(self, in_channels, out_channels256, rates[6, 12, 18]): super().__init__() self.branches nn.ModuleList() # 1x1 卷积分支 self.branches.append(nn.Sequential( nn.Conv2d(in_channels, out_channels, 1, biasFalse), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue) )) # 不同空洞率的 3x3 卷积分支 for r in rates: self.branches.append(nn.Sequential( nn.Conv2d(in_channels, out_channels, 3, paddingr, dilationr, biasFalse), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue) )) # 全局池化分支 self.global_pool nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Conv2d(in_channels, out_channels, 1, biasFalse), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue) ) self.project nn.Sequential( nn.Conv2d(out_channels * (len(rates) 2), out_channels, 1, biasFalse), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue), nn.Dropout(0.5) ) def forward(self, x): size x.shape[-2:] feats [branch(x) for branch in self.branches] gp self.global_pool(x) gp nn.functional.interpolate(gp, sizesize, modebilinear, align_cornersFalse) feats.append(gp) out torch.cat(feats, dim1) return self.project(out)paddingr和dilationr必须相等这样输出尺寸才和输入一致。Dropout(0.5)是防止过拟合如果数据集小可以调到 0.3。全局池化分支最后要插值回原尺寸再拼接align_cornersFalse是 PyTorch 分割任务里的常见设置和True的区别在边界像素上用错会导致边缘预测偏移。3.3 分类头与输出层ASPP 输出后接一个 1x1 卷积把通道数映射到类别数。Cityscapes 训练用 19 类如果直接用 30 类原始标签这里改成 30。class DeepLabV3(nn.Module): def __init__(self, backboneresnet101, num_classes19): super().__init__() self.backbone, in_ch build_backbone(backbone) self.aspp ASPP(in_ch, 256) self.classifier nn.Conv2d(256, num_classes, 1) def forward(self, x): size x.shape[-2:] feat self.backbone(x) feat self.aspp(feat) out self.classifier(feat) # 上采样回输入尺寸 out nn.functional.interpolate(out, sizesize, modebilinear, align_cornersFalse) return out输出上采样回原图尺寸是为了和标签算 loss 时维度对齐。如果显存不够可以在低分辨率上算 loss 再上采样但代码里通常直接插值。num_classes这个参数一定要和数据集类别数一致否则交叉熵会报 index 越界。4. 训练循环与损失函数从交叉熵到忽略标签的处理4.1 损失函数与忽略标签Cityscapes 里有些像素是无效的比如 ego vehicle 或者未标注区域标签 ID 是 255。交叉熵要设ignore_index255否则这些像素会参与梯度计算拉偏模型。import torch.nn as nn criterion nn.CrossEntropyLoss(ignore_index255)如果类别不均衡严重可以加类别权重但 Cityscapes 的 19 类分布还算均匀一般不加也能收敛。ignore_index这个参数是必须的忘了设的话 loss 会莫名其妙偏高而且验证集指标上不去。4.2 优化器与学习率策略常见做法是 SGD 加 poly 衰减初始学习率 0.01动量 0.9权重衰减 1e-4。poly 衰减的公式是lr base_lr * (1 - iter / max_iter) ** 0.9。import torch.optim as optim optimizer optim.SGD(model.parameters(), lr0.01, momentum0.9, weight_decay1e-4) def poly_lr(base_lr, iter, max_iter, power0.9): return base_lr * (1 - iter / max_iter) ** power # 训练循环里每个 iter 更新 for it, (imgs, labels) in enumerate(loader): lr poly_lr(0.01, it, max_iter40000) for param_group in optimizer.param_groups: param_group[lr] lr imgs, labels imgs.cuda(), labels.cuda() outputs model(imgs) loss criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step()max_iter设 40000 是 Cityscapes 上的常见值batch size 8 的话大概跑几十个 epoch。poly 衰减比 step 衰减更平滑后期学习率降得慢有利于收敛到更优解。如果 loss 震荡厉害先把学习率降到 0.005 试试。4.3 验证指标与模型保存验证时算 mIoU按类别求交并比再平均。保存模型建议按 mIoU 最高的存而不是按 loss 最低因为 loss 和分割质量不完全正相关。def compute_iou(pred, label, num_classes19, ignore255): pred pred.argmax(dim1) mask label ! ignore pred pred[mask] label label[mask] ious [] for c in range(num_classes): inter ((pred c) (label c)).sum().item() union ((pred c) | (label c)).sum().item() if union 0: ious.append(inter / union) return sum(ious) / len(ious) if ious else 0.0argmax(dim1)是在类别维度上取最大得到每个像素的预测类别。mask过滤掉忽略标签。如果某个类别在验证集里没出现union为 0跳过不计入平均这是常见处理方式。保存模型时用torch.save(model.state_dict(), best.pth)只存权重加载时先建模型再load_state_dict。5. 避坑与排查训练不收敛、显存溢出、指标异常的常见原因5.1 现象loss 一直不降停在 2.9 左右原因通常是标签映射错了。如果你用的是labelTrainIds.png但代码按labelIds的 19 类去读类别 ID 对不上模型学不到有效信息。解决方法是确认读的是labelIds还是labelTrainIds两者选其一并保证num_classes和映射逻辑一致。5.2 现象CUDA out of memory原因可能是 batch size 太大、裁剪尺寸太大或者没释放中间变量。先把 batch size 降到 4裁剪尺寸降到 256x512 试试。如果还不行检查是不是在验证时忘了torch.no_grad()导致计算图一直累积。加上with torch.no_grad():能省不少显存。5.3 现象mIoU 卡在 0.3 上不去原因可能是学习率太大导致震荡或者 ASPP 的 rates 和输入尺寸不匹配。先看训练 loss 是否正常下降如果 loss 降但 mIoU 不涨多半是过拟合或者验证集预处理和训练不一致。检查验证时有没有做同样的归一化均值和方差是否一致。5.4 现象预测结果全是同一类原因通常是类别权重严重失衡或者最后一层卷积初始化有问题。检查classifier的权重初始化默认 PyTorch 的初始化一般没问题但如果自己改了要确认。另外看看训练数据里是不是某一类占了绝大多数如果是考虑加类别权重或者重采样。5.5 现象训练速度特别慢原因可能是数据加载成了瓶颈。num_workers设成 4 或 8pin_memoryTrue能明显加快。如果用的是机械硬盘数据读取慢是硬伤换 SSD 或者把数据预加载到内存里。另外确认cudnn.benchmark True有没有开开了能自动选最优卷积算法。6. 进阶技巧用预训练权重加速收敛与推理部署6.1 加载预训练权重做迁移学习如果不想从零训练可以加载在 Cityscapes 上已经训好的权重只微调最后几层。常见做法是冻结主干的前面几层只训 ASPP 和分类头。model DeepLabV3(backboneresnet101, num_classes19) state torch.load(deeplabv3_cityscapes.pth, map_locationcpu) model.load_state_dict(state, strictFalse) # 冻结主干前 5 层 for name, param in model.backbone.named_parameters(): if layer1 in name or layer2 in name: param.requires_grad False optimizer optim.SGD(filter(lambda p: p.requires_grad, model.parameters()), lr0.001, momentum0.9)strictFalse允许部分权重不匹配适合你改了分类头类别数的情况。冻结浅层是因为浅层特征通用性强没必要重新学。学习率调小到 0.001因为预训练权重已经不错了大步长容易破坏。6.2 推理脚本与单张图像预测推理时要注意模型切到eval()模式并且用no_grad包住。model.eval() with torch.no_grad(): img Image.open(test.png).convert(RGB) img_t T.functional.to_tensor(img) img_t T.functional.normalize(img_t, mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) img_t img_t.unsqueeze(0).cuda() out model(img_t) pred out.argmax(dim1).squeeze().cpu().numpy()unsqueeze(0)是加 batch 维度因为模型要求输入是 4D。argmax后得到每个像素的类别索引可以映射成颜色可视化。如果推理结果和训练时差距大检查eval()有没有调BatchNorm 在训练和推理时的行为不一样。6.3 导出 ONNX 做部署如果需要部署到非 PyTorch 环境可以导出 ONNX。dummy torch.randn(1, 3, 512, 1024).cuda() torch.onnx.export(model, dummy, deeplabv3.onnx, input_names[input], output_names[output], opset_version11)opset_version11兼容性较好dummy的尺寸要和实际推理一致。导出后可以用 onnxruntime 验证输出是否和 PyTorch 一致差异在 1e-3 以内算正常。从那以后我每次跑分割任务都强制先验证一遍数据加载的标签映射和归一化参数这两处出错最隐蔽也最耗时间。希望帮到你。本文还有配套的精品资源点击获取