ARTICLE DETAIL

资讯详情

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

PSMNet复现全流程踩坑记录:从环境配置到KITTI训练

PSMNet复现全流程踩坑记录:从环境配置到KITTI训练 搞立体匹配方向的同学一定绕不开PSMNet这个经典模型。我在两周内从零复现它期间踩了环境配置、数据读取、显存爆炸、指标计算等各种坑最后总算把它在KITTI上完整跑通。这篇文章把我整个踩坑过程从环境搭建到数据集训练全部梳理成一份可直接上手的操作记录里面包含了我实际用的命令、代码片段和排查思路。无论你是刚入门深度学习想找个经典模型练手还是已经在做立体匹配相关课题这份记录应该都能帮你少走不少弯路。1. 先看一遍PSMNet再动手核心思路与复现价值1.1 为什么要复现这个2018年的老模型虽然PSMNet是2018年的工作但直到今天KITTI立体匹配排行榜上依然能看到它的身影很多后续方法也都拿它当baseline对比。对新手来说复现PSMNet能一次性接触到立体匹配里最核心的几个套路特征提取、代价体构建、3D卷积聚合、soft argmin回归视差。这套pipeline搞明白了再看RAFT-Stereo、IGEV这类新模型会轻松很多。我自己复盘整个过程最大的感受是复现经典论文的价值不在于把指标刷多高而是把“读论文”变成“跑通代码再回头看论文”很多抽象概念在亲手实现之后才有实感。比如cost volume这东西看论文时觉得不就是特征拼接嘛真正自己去写那个循环shift左特征和右特征的代码时才明白为什么它是整个模型里最吃显存的部分。1.2 PSMNet的核心结构拆解PSMNet的大致流程是这样的左右两幅图先经过一个共享权重的CNN提取特征输出分辨率是原图的四分之一通道数一般是32。然后特征进入空间金字塔池化SPP模块用不同尺寸的平均池化获取多尺度上下文信息。接下来是关键的cost volume构建对于每个视差值d把左特征和右特征在宽度方向上做平移后拼接得到形状为[B, 2C, D, H/4, W/4]的代价体。之后用堆叠的3D卷积对这堆代价体做正则化最后在视差维度上做softmax加权求和得到亚像素精度的视差图。这里有几个复现时必须注意的细节。SPP模块的输出尺寸对齐是个隐藏坑不同池化核大小算出来的特征图尺寸可能差一个像素直接上采样回原尺寸再接1x1卷积容易出问题。我的建议是把池化层的stride和kernel size都设置成能整除的数值或者干脆用AdaptiveAvgPool2d省心很多。另外就是maxdisp这个超参数。KITTI数据集的视差范围通常不超过192像素所以大多数实现默认maxdisp192。这个值决定了cost volume在视差维度上的尺寸直接影响显存占用。复现时先按192跑万一你的应用场景视差范围特别大再调整也不迟。1.3 三个“推理正确但跑不通”的设计细节第一个是视差回归。PSMNet不是直接输出一个标量视差而是对cost volume在视差维度上做softmax得到概率分布然后对每个视差级别做加权求和。这个操作在论文里叫soft argmin。实现起来很简单几行代码就能搞定但有个容易出错的地方softmax的logits如果不做缩放概率分布可能会过于尖锐或平坦导致训练初期梯度异常。有些实现会在softmax之前除以一个温度系数实测下来更稳。第二个是中间监督。原版PSMNet的3D CNN部分包含堆叠的hourglass结构会在不同阶段输出视差估计并加上损失。复现时如果仓库代码里有这种多尺度监督别急着删掉它对收敛速度有明显帮助。我第一次复现时把中间监督全去掉了结果训练到200个epoch时loss还降不到理想值加回去之后明显好转。第三个是输出分辨率的问题。网络最后输出的视差图是原始分辨率的四分之一计算损失前需要上采样回原图大小和ground truth对齐。这个上采样一般用双线性插值就能搞定但要注意如果你后面要做左右一致性检查或者后处理最好在四分之一分辨率上先做再上采样不然边缘会糊。2. 环境配置的深坑Python、PyTorch、CUDA与扩展模块2.1 版本组合选择不要装最新要装稳定环境配置是第一个让我抓狂的环节。PSMNet这个项目本身不算复杂但它依赖的PyTorch版本和你显卡驱动之间必须匹配否则会出现各种莫名奇妙的报错。我最终选定的组合是Python 3.8 PyTorch 1.13.1 CUDA 11.7跑得非常稳。在动手装环境之前第一步一定是先看自己显卡驱动支持的CUDA版本用nvidia-smi查看右上角的CUDA Version。注意这里显示的版本是驱动支持的最高CUDA版本不代表你机器上已经装了对应CUDA工具包。PyTorch安装时选择CUDA 11.7还是11.8只要不超过驱动支持的最高版本就行。如果你用的是Anaconda创建环境的命令大概是这样的conda create -n psmnet python3.8 -y conda activate psmnet pip install torch1.13.1cu117 torchvision0.14.1cu117 --extra-index-url https://download.pytorch.org/whl/cu117装完之后用一行命令验证GPU是否可用python -c import torch; print(torch.__version__, torch.cuda.is_available(), torch.cuda.get_device_name(0))我当时遇到过一种情况torch.cuda.is_available()返回True但真跑模型时报CUDA error: no kernel image is available for execution on the device。这通常是PyTorch版本和显卡架构不匹配导致的比如太老的显卡装了太新的PyTorch。遇到这种问题就只能换一个兼容你显卡架构的PyTorch版本没有别的捷径。2.2 CUDA扩展编译gcc和nvcc的版本战争市面上很多PSMNet的开源实现除了纯PyTorch代码之外还需要编译一些自定义的CUDA扩展比如resample2d、channelnorm之类的模块。这些扩展在训练和推理时用于处理特征图对齐功能上用纯PyTorch也能实现但原版代码为了追求效率选择了CUDA实现。编译这类扩展最坑的就是gcc版本和nvcc不兼容。我自己遇到过一个经典报错error: #error -- unsupported GNU version! gcc versions later than 8 are not supported!这个报错的意思是你机器上的gcc版本太新超出了当前CUDA版本支持的范围。解决办法有几种如果系统里装了多个gcc版本可以用环境变量指定一个旧版本export CUDAHOSTCXX/usr/bin/g-8如果是conda环境可以装一个低版本的gccconda install gxx_linux-647如果源码是用setup.py安装的也可以临时改一下编译参数另一个常见报错是fatal error: cuda_runtime.h: No such file or directory这是因为编译时找不到CUDA头文件。解决办法是设置CUDA_HOME环境变量export CUDA_HOME/usr/local/cuda export PATH$CUDA_HOME/bin:$PATH export LD_LIBRARY_PATH$CUDA_HOME/lib64:$LD_LIBRARY_PATH2.3 其余依赖包清单除了PyTorch之外还需要几个常用库opencv用于图像读取tensorboardX或tensorboard用于可视化训练曲线matplotlib用于画图tqdm用于显示进度条。一次性装齐可以省很多事pip install opencv-python tensorboardX matplotlib tqdm这里有个小提醒opencv-python和opencv-contrib-python不要同时装会冲突。读KITTI的16位视差图用普通版opencv就够了。IDE方面我用的VSCode配合Remote SSH直接在服务器上改代码、看日志、可视化都挺方便。如果你也是用VSCode记得装Python扩展和Pylance调试断点功能对排查数据加载问题很有帮助。3. KITTI数据集的下载、读取与加载器实现3.1 数据集的目录结构和训练/测试划分KITTI立体匹配数据集分2012和2015两版我复现时用的是KITTI 2015。下载完解压之后目录结构是这样的data_scene_flow/ training/ image_2/ 左目彩色图像 image_3/ 右目彩色图像 disp_occ_0/ 带遮挡标注的视差图 disp_noc_0/ 非遮挡区域的视差图 testing/ image_2/ image_3/training目录下共有200对图像testing目录下也有200对左右。图像分辨率大约是1242x375整体偏宽幅。没有官方指定的训练集和验证集划分常见做法是把training里面的图像随机分成160对训练、40对验证或者直接用一个固定随机种子划分方便对比实验结果。KITTI 2012的目录结构略有不同用的是colored_0和colored_1来存放左右图但处理逻辑基本一致。如果你需要对照排行榜上的历史结果2012和2015都有完整的排名可以挑一个基准数据来验证自己的复现效果。3.2 视差图的正确读取方式这是整个复现过程里最容易忽略但影响最大的一个细节。KITTI的视差图是16位PNG格式真实视差值需要除以256才能得到浮点数。很多人第一次读的时候直接用cv2.imread默认参数读出来是8位图所有视差值都变成了0或者255训练loss直接崩掉或者变成奇怪的值。正确读取代码如下import cv2 import numpy as np def read_disp(path): disp cv2.imread(path, cv2.IMREAD_UNCHANGED) disp disp.astype(np.float32) / 256.0 return dispdisp_occ_0和disp_noc_0的区别在于前者包含遮挡区域的标注后者只标注非遮挡区域遮挡区域像素值为0。训练时一般用disp_occ_0配合一个mask把视差值等于0的无效像素过滤掉。评估时如果想和排行榜对齐需要明确自己用的是非遮挡还是全部区域指标这两个数值差不少。我还踩过一个隐蔽的坑有的读取代码用PIL.Image.open读取16位PNG但PIL默认可能把它当成I;16模式直接.astype(np.float32)不做除以256导致loss异常大。所以不管用什么库读一定记得除以256并且做一次可视化确认。3.3 Dataset与DataLoader实现要点Dataset类的核心逻辑其实很简单根据索引拿到对应的左图、右图、视差图做一些预处理最后返回三个张量。我实现时大概是这样的结构class KITTIDataset(Dataset): def __init__(self, root, file_list, transformNone): self.left_paths file_list self.right_paths [p.replace(image_2, image_3) for p in file_list] self.disp_paths [p.replace(image_2, disp_occ_0).replace(.png, _disp_occ_0.png) for p in file_list] self.transform transform def __getitem__(self, idx): left cv2.imread(self.left_paths[idx]) right cv2.imread(self.right_paths[idx]) disp read_disp(self.disp_paths[idx]) # 这里根据文件名映射需要仔细核对KITTI的文件名后缀有点绕 # 比如 000000_10.png 对应的视差图是 000000_10_disp_occ_0.png if self.transform: left, right, disp self.transform(left, right, disp) return left, right, disp文件名的映射是另一个容易踩坑的地方。KITTI 2015的左右图和视差图文件名不是简单的前缀替换而是类似000000_10.png对应000000_10_disp_occ_0.png这种多了一段后缀。我一开始想当然地用replace(.png, .png)结果一堆文件找不到排查半天才发现是命名规则的问题。一个更省事的方式是直接遍历目录把文件名和对应的左右图、视差图路径存成列表。因为KITTI的数据集不大没必要在__getitem__里做复杂的路径拼接启动时把这个映射关系一次性构建好就行。DataLoader部分有几个效率相关的设置train_loader DataLoader( train_dataset, batch_size8, shuffleTrue, num_workers4, pin_memoryTrue, drop_lastTrue )num_workers设置成4或者8可以显著提升数据加载速度但也不是越大越好。我曾在服务器上把num_workers设成16结果训练启动时直接段错误后来降到8就好了。pin_memoryTrue在GPU训练时能减少数据传输时间建议开启。如果你发现训练过程中GPU利用率经常掉到很低多半是数据加载跟不上了优先检查这两项。3.4 数据增强中的几何变换坑训练时做数据增强有利于提高泛化性但立体匹配任务的数据增强比普通分类任务麻烦因为左右图和视差图三者必须同步变换而且有些变换会改变视差语义。我踩过最深的一个坑是水平翻转。如果只是简单地把左右图和视差图同时水平翻转其实破坏了左右图的对应关系。正确的做法有两种进行水平翻转时将左右图交换同时将视差图水平翻转视差值本身不变或者干脆不做水平翻转只做颜色抖动、亮度调整、高斯噪声这类不会改变几何关系的增强我后来偷懒选了第二种只在颜色和亮度上做扰动效果还不错。如果你想做水平翻转强烈建议先用一对样本做可视化确认翻转后的左右图和视差图语义是对得上的再把它加进训练流程。垂直翻转倒是安全很多因为视差只跟水平方向有关垂直翻转对视差值没有影响左右图同时垂直翻转即可。随机裁剪也是常用增强手段裁剪左右图和视差图的同一区域视差值不变。常见的裁剪尺寸是256x512或320x736裁剪太小会丢失上下文信息太大又会让显存压力剧增。4. 训练流程、超参与显存优化实操4.1 损失函数与评估指标怎么算PSMNet的损失函数用的是Smooth L1 Loss也就是Huber Loss对离群点没那么敏感。实现时要注意用mask过滤掉无效视差像素import torch.nn.functional as F def smooth_l1_loss(pred, gt, mask): loss F.smooth_l1_loss(pred, gt, reductionnone) loss loss[mask].mean() return loss评估指标方面KITTI最常用的两个指标是EPE和D1。EPE是平均端点误差即预测视差和真实视差的平均绝对像素差。D1是误匹配率统计误差大于3像素且相对误差超过5%的像素比例。计算D1的代码大概长这样def compute_metrics(pred, gt, mask): pred, gt pred[mask], gt[mask] epe (pred - gt).abs() err ((epe 3.0) (epe / gt 0.05)).float().mean() return epe.mean().item(), err.item() * 100如果你在验证集上看到D1在3%到5%之间EPE在1像素左右基本说明复现是成功的。原论文报告的结果会比随机验证集划分好一些因为官方评测集和训练集分布更一致不用强行追求完全一样的数字。4.2 超参数设置背后的逻辑我训练时用的关键超参数如下参数数值说明batch_size8受限于显存24GB显卡可以到16maxdisp192KITTI视差范围上限learning_rate0.001Adam优化器初始学习率epochs400训练轮数可视收敛情况调整crop_size256x512随机裁剪尺寸weight_decay0.0001L2正则化系数学习率策略上原论文是每200个epoch衰减10倍。我在实际训练中发现400个epoch的进度下前200个epoch用0.001之后切到0.0001模型能比较好地收敛。这个策略不是最优的但省心稳定不至于在中途因为学习率太大而loss发散。优化器我用的Adam比SGD收敛快不少。对复现任务来说Adam是最省事的选择。如果你的显卡显存紧张batch_size被迫设得很小可以考虑用梯度累积来模拟较大的batch size效果更接近原论文的设定。4.3 显存爆炸的处理手段PSMNet最大的特点就是吃显存。我粗略算过裁剪尺寸256x512特征图四分之一分辨率就是64x128通道数32maxdisp为192时cost volume本身的形状是[B, 64, 192, 64, 128]。光这一个张量batch size为8时就要占用约100GB不对我重算一下。每对特征的cost volume不考虑batch大小是2倍通道数左右特征拼接所以是64通道乘以192个视差级别再乘以64x128的空间分辨率。所以单个样本的cost volume约是64x192x64x128 1亿个float元素换算成内存大概是400MB左右。这是构建出来的代价体还没算3D卷积的中间结果。batch size为8时光这层就到3.2GB了再加上后续3D卷积的特征图24GB显存也就勉勉强强。所以我复现时做了一个取舍batch size从论文常用的16降到8裁剪尺寸从320x736降到256x512。如果你只有12GB显存建议batch size设为4裁剪尺寸改成192x384或者把maxdisp降到160先验证流程。另外有个小技巧不需要高频执行torch.cuda.empty_cache()这命令看似释放了显存实际上会降低性能。除非你在同一个进程里反复加载和释放大模型否则让PyTorch自己管理就好。4.4 训练日志与模型保存训练时建议定期打印loss和验证指标同时用tensorboard记录训练曲线。我实际跑出来的日志大概长这样epoch 080, lr 0.001, loss 0.214, epe 1.38, d1 10.4% epoch 160, lr 0.001, loss 0.177, epe 0.98, d1 6.1% epoch 240, lr 0.0001, loss 0.164, epe 0.89, d1 4.8% epoch 320, lr 0.0001, loss 0.158, epe 0.84, d1 4.2%如果你看到loss一直在降但验证集指标不动可能是过拟合了需要加一些正则化或者数据增强。如果loss直接变成NaN大概率是数值稳定性的问题可以把输入图像归一化到0到1之间或者降低学习率。模型保存务必使用checkpoint方式把model、optimizer、scheduler、epoch等信息都存下来方便断点续训torch.save({ epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), scheduler_state_dict: scheduler.state_dict(), }, checkpoint_latest.pth)加载checkpoint时要注意模型是否用了DataParallel或DistributedDataParallel如果保存时带module.前缀加载到普通模型上会报key不匹配的错需要把前缀去掉state_dict torch.load(checkpoint_latest.pth)[model_state_dict] from collections import OrderedDict new_state_dict OrderedDict() for k, v in state_dict.items(): if k.startswith(module.): k k[7:] new_state_dict[k] v model.load_state_dict(new_state_dict)5. 常见问题排查与避坑技巧实录5.1 报错速查表整个复现过程中我遇到过的典型报错整理成了一张表方便快速定位问题现象可能原因解决办法CUDA out of memorycost volume或3D卷积占用显存过大降低batch_size、裁剪尺寸、maxdisptorch.cuda.is_available()为FalsePyTorch版本与CUDA不匹配重装PyTorch确认nvidia-smi驱动版本编译扩展时报gcc版本不支持nvcc与gcc版本不匹配设置CUDAHOSTCXX或安装旧版本gcc训练loss快速降为0或固定值视差图读取错误没除以256检查视差图读取代码可视化验证eval指标和训练差距过大忘了model.eval()或torch.no_grad()评估前加model.eval()和with torch.no_grad()验证集D1指标一直在20%以上验证集划分或mask处理有问题检查mask过滤是否正确确认评估代码DataLoader启动时报段错误num_workers过大降低num_workers尝试设为0加载checkpoint时key不匹配DataParallel前缀问题去掉module.前缀训练时GPU利用率低数据加载是瓶颈增加num_workers开启pin_memory5.2 我反复踩的几个隐性坑除了上面表格里的问题还有几个隐性但影响很大的坑。第一个是验证时忘记把BatchNorm层切换到eval模式。BatchNorm在训练和推理时的行为不一样如果忘了model.eval()验证集指标会高得离谱而且每次验证结果都不一样。我当时就因为这个以为模型训练崩了排查了半天才发现是这个小问题。第二个坑是验证集划分。KITTI training只有200对图像如果你在__getitem__里直接用整个目录做训练没有单独留验证集最终打印的“验证”指标其实是在训练集上算的看起来很好但实际泛化能力未知。正确做法是先把文件列表切分出训练和验证两部分再分别构造Dataset。第三个是评估时的mask选择。如果用了disp_occ_0做训练但评估时也直接用disp_occ_0并只过滤0值像素那么D1指标会偏高因为这些像素本来不需要被评估。KITTI官方评估通常分为非遮挡区域和全部区域两个指标。做复现对比时统一一个mask逻辑千万别今天用noc明天用occ最后得到的结果完全没有可比性。5.3 一些加快实验节奏的调试技巧最后分享几个我后来逐步养成的实验习惯。先跑通小规模再上全量。我第一次复现时直接把全部200对图像、400个epoch一次跑满结果第三天发现代码里有个bug浪费了大量时间。后来我学乖了先用10对图像、10个epoch跑一遍只要能正常出loss、指标在下滑再切回全量训练。这个冒烟测试看起来多花半小时实际上能省好几天。固定随机种子。PyTorch的DataLoader、模型初始化、数据增强都有随机性如果不固定种子可能连续两次训练得到的结果差异很大干扰你判断代码改动是否有效。建议在训练脚本开头加上import random import numpy as np import torch def set_seed(seed): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed)及时可视化视差图。训练到一半把预测的视差图保存下来用伪色彩图看一眼。如果预测图边缘有大量黑色条纹可能是mask或上采样处理不对。如果整张图颜色很淡可能是soft argmin后视差范围偏小。可视化能帮你快速发现很多loss指标看不出来的问题。把实验配置写进文件名或json配置里。模型结构、超参数、数据增强策略、checkpoint路径这些信息最好和日志、模型权重一起保存下来。我自己就吃过亏两周后回看一个模型的权重完全想不起来当时用了什么配置只能重新跑实验对比。我在实际使用中发现复现PSMNet最大的收获倒不是最终指标有多好看而是整个pipeline里每个环节都亲手摸了一遍。环境配置让你学会管理深度学习环境数据读取让你理解KITTI的存储格式训练调试让你真正明白损失函数和评估指标的关系。等这套流程跑通后面想换新的立体匹配模型基本上只需要改网络结构部分其他环节都能复用。最后再分享一个小技巧把KITTI的读取、评估、可视化写成独立模块一开始就同时跑起来调试网络结构时会轻松非常多。
返回列表