ARTICLE DETAIL

资讯详情

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

PSMNet复现避坑指南:KITTI双目视差估计从环境配置到训练评估全攻略

PSMNet复现避坑指南:KITTI双目视差估计从环境配置到训练评估全攻略 从拿到一份论文代码到真正复现出论文效果中间隔着的不是代码本身而是一长串的坑。如果模型是PSMNet这种2018年的经典老代码那坑的密集程度会成倍上升。我是在一个需要做双目视差估计的项目里决定用PSMNet做baseline的一开始以为“clone仓库、装好依赖、跑一跑就行”结果从环境配置到KITTI数据集训练整整折腾了一周。这篇就把整个过程中踩过的坑、排查思路、最终能稳定训练的最小配置全部记录下来给后面想复现PSMNet或者类似老项目的朋友做一个参考。这篇内容适合几类人看第一次接触立体匹配/视差估计、想跑通PSMNet作为baseline的研究生或工程师已经跑通过但被KITTI数据集的预处理和评估指标绕晕的人以及在PyTorch老版本、老代码迁移到新环境时经常被版本兼容问题折磨的玩家。文章不会止于“能跑”会一直讲到loss曲线正常下降、能输出像样的视差图、能算KITTI的D1指标为止。1. 复现PSMNet之前先把这些基础认知对齐1.1 PSMNet到底在解决什么问题PSMNet全称是Pyramid Stereo Matching Network2018年CVPR的文章作者是Jia-Ren Chang和Yong-Sheng Chen。它的任务是双目立体匹配输入是左右两张经过了极线校正的RGB图像输出是一张视差图图上每个像素的值表示该点在左图和右图之间的水平位移。根据视差、基线长度和焦距可以进一步换算出深度。当时PSMNet能在KITTI 2012和KITTI 2015榜单上拿到很好的成绩核心是两个设计一是空间金字塔池化模块SPP用来在不同尺度上聚合全局上下文信息解决弱纹理区域的匹配歧义二是堆叠沙漏结构的3D CNN对代价体进行正则化让视差估计在边缘和遮挡区域更干净。后面大量工作都是在这两个思路上做改进的比如GwcNet引入分组相关CFNet做级联细化但PSMNet作为baseline的地位很稳。所以在复现之前要有一个明确认知这不是“跑通一个demo”就完事的项目而是一个完整的训练-验证-评估链路。你要准备的不只是模型代码还有数据集的下载与解析、训练pipeline的调整、以及评估指标的对齐。1.2 代码很“老”这件事要提前做好心理建设PSMNet官方开源代码是基于PyTorch 0.4-era写的有很多那个年代特有的写法。比如Python 2风格的print语句、xrange、字符串格式化方式、.data的到处使用等。跑在新版本PyTorch上大概率直接在import阶段就报错。我的建议是不要把“让老代码原封不动跑起来”当成目标而是把“把老代码迁移到新环境”当成目标。这本质上是一个移植工程等于要动不少代码。好在PSMNet的代码量不大核心文件就那么几个models/psmnet.py是模型结构datasets/kitti_dataset.py是KITTI数据加载逻辑main.py是训练入口动起来比想象中可控。1.3 硬件门槛先摸清楚PSMNet的显存占用是天生的硬指标。它要构建一个4D代价体维度是[D, H/2, W/2, 2C]其中D是最大视差搜索范围默认192C是特征通道数默认32。以常见的256x512裁剪输入为例代价体本身就已经非常占显存再加上3D CNN的中间特征图batch size稍微调大一点显存就会爆炸。我用RTX 4090 24GB实测batch size设10到12是能跑的。如果用更小的显卡比如8GB显存就需要把crop尺寸缩到192x384batch size缩到4代价是训练效果会打折扣。所以复现之前先确认手头的GPU这直接决定了后面一系列超参怎么设。2. 环境配置的版本地狱我最终采用的组合与排查链路2.1 为什么老项目在“环境”上最容易脱层皮PyTorch的接口演进在过去几年非常剧烈尤其是从0.4到1.x再到2.x很多API要么改名、要么行为变了。PSMNet依赖的旧API和现代CUDAToolkit/PyTorch之间并没有稳定的兼容关系所以环境配置成了复现路上的第一道坎。我最初直接装了最新的PyTorch 2.1 CUDA 12.1结果遇到一堆问题首先是编译模型时torch.nn.functional里某个接口行为变了其次是th模块根本不存在再次是有些算子在新版本上数值表现不稳定。折腾一圈后我意识到问题不在某个具体报错而是“新版本环境 老代码”这个组合本身就充满了不确定性。2.2 我最后选定的版本组合踩了一圈之后我最终稳定跑通的组合是这样的组件推荐版本备注Python3.7老代码兼容性最好CUDA Toolkit10.1 / 11.3取决于显卡驱动Ampere架构建议11.3cuDNN7.6.5 / 8.2与CUDA版本匹配即可PyTorch1.6.0 / 1.10.01.6开始支持AMP但PSMNet建议直接用FP32GCC7.5编译扩展模块够用numpy1.19.51.20会对某些老代码报错opencv-python4.5.x读取KITTI图片足够这个组合的核心逻辑是PyTorch 1.6到1.10之间保留了大部分旧API的兼容层而Python 3.7则同时满足现代语法和老代码的习惯。千万不要追求版本新这个项目的目标不是体验新特性而是稳定复现。如果你的显卡是RTX 30系列CUDA Toolkit不能低于11.0否则PyTorch的CUDA算子根本无法在你的GPU上运行。RTX 40系列的话驱动版本足够的话CUDA 11.3也可以支持。2.3 创建环境的完整命令我建议直接用conda管理环境省心。关键命令如下conda create -n psmnet python3.7 conda activate psmnet conda install pytorch1.6.0 torchvision0.7.0 cudatoolkit10.1 -c pytorch pip install opencv-python4.5.5.64 pip install tensorboard pip install numpy1.19.5 pip install scikit-image装完之后先跑一个最基础的smoke test确认PyTorch能看到GPUimport torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0)) x torch.randn(4, 32, 128, 256).cuda() y torch.conv3d(x, torch.randn(64, 32, 3, 3, 3).cuda(), padding1) print(y.shape)这段能跑通说明CUDA、cuDNN、PyTorch三者的底层链路没问题。这个验证非常值得做因为很多人在还没确认底层环境的时候就急着跑模型遇到报错会分不清是环境问题还是代码问题。2.4 老代码迁移到新环境的几个必改位置即使版本组合选对了代码本身还是要动刀。我在迁移过程里改得最多的是这几处xrange全部替换为range。这个无脑全局替换即可。print语句加括号。原代码里有多处print xxx的形式Python 3下直接报错。去掉.data。虽然.data在PyTorch 1.x还保留但PyTorch 1.6以后会警告而且在新版本上.data拿到的Tensor不会自动加入计算图容易埋坑。直接把a.data改成a.detach()更安全。torch.squeeze(a, 1)这类调用如果在新版本上报错检查是否需要改成torch.squeeze(a, dim1)因为不同版本的函数签名变化比较坑。原来代码里如果有model torch.nn.DataParallel(model)的写法要确保在创建优化器之后再wrap否则model.parameters()拿不到完整参数。这个坑在单卡训练时不会暴露一旦想用多卡就炸。3. KITTI数据集的下载、目录结构与预处理陷阱3.1 KITTI 2015和KITTI 2012怎么选KITTI数据集是双目视觉领域绕不开的基准。做PSMNet复现时可以选KITTI 2015和KITTI 2012两套数据集它们的内容和目标不完全一样对比项KITTI 2015KITTI 2012全称KITTI Stereo 2015 / Scene FlowKITTI Stereo 2012图像内容道路场景包含动态车辆道路场景静态为主真值类型视差 光流 场景流视差常用于测试PSMNet更常用论文报告该榜同样可以跑下载文件data_scene_flow.zipdata_stereo_flow.zipPSMNet论文里报告了两个榜的结果但绝大多数复现项目更关注KITTI 2015因为它的真值更完整且是后来的主要benchmark。我的建议是直接下KITTI 2015也就是data_scene_flow.zip训练和评估一整套流程都在这套数据上走通。KITTI官网需要注册后才能下载。下载完成以后你会拿到一个很大的zip解压出来的目录结构大致是这样的data_scene_flow/ ├── training/ │ ├── image_2/ # 左目彩色图 │ ├── image_3/ # 右目彩色图 │ ├── disp_occ_0/ # 左目视差真值(带遮挡标记) │ └── calib_cam_to_cam.txt └── testing/ ├── image_2/ ├── image_3/ └── calib_cam_to_cam.txt注意image_2是左目image_3是右目这个和某些人的直觉是反的。我一开始就搞反过结果训练出来的模型左图右图在代价体构建时全乱了loss曲线非常难看。3.2 视差真值的读取一个最容易翻车的细节KITTI的视差真值不是普通的灰度图。它用PNG文件存储float类型的视差值具体做法是把真实的视差值乘以256后存成一个整型PNG。所以在读取时必须除以256才能得到真正的视差。一个标准的读取方式是这样的def read_disp(filename): disp cv2.imread(filename, cv2.IMREAD_UNCHANGED) disp disp.astype(np.float32) / 256.0 return disp如果直接用cv2.imread(filename, cv2.IMREAD_GRAYSCALE)读取你会拿到一个被截断的灰度图视差范围完全不对训练出来的模型也会是一个废物。另外视差真值中像素值为0的地方代表无效区域或遮挡区域训练loss计算时需要把这里的贡献mask掉否则模型会被无效像素带到沟里。3.3 训练集和验证集的划分KITTI官方并没有为训练算法提供一个明确的train/val划分所以复现PSMNet时通常采用原项目或者社区约定俗成的划分方式。一个很常见的做法是把training目录下0到159的160张图像作为训练集160到199的40张图像作为验证集。原版PSMNet就是这么处理的。划分的逻辑很简单用连续的序号切分保证两个集合的场景不重叠也便于复现论文里的数据。如果你自己写dataloader记得和这个划分保持一致否则你看到的loss曲线和论文里对不上会以为自己复现出了bug。3.4 图像的尺寸不一致问题KITTI 2015里的图像尺寸并非完全统一虽然大部分是1242x375但边缘有一些是1238x374之类的。PSMNet在训练和推理时通常会把输入resize到固定的尺寸。原论文和原代码里用的是crop到256x512或类似尺寸。我实际使用时在dataloader里直接random_crop到[256, 512]训练完后再对整图推理效果是OK的。要特别留神的是resize、crop这些操作一定要同时作用于左图、右图和视差真值否则左右图对不齐代价体直接废掉。写数据预处理的时候最好封装成一个函数同时处理三种输入避免人为疏忽。4. 训练阶段的最大拦路虎显存规划、loss分析与超参调整4.1 显存不够用时的降级方案PSMNet训练时最痛苦的问题就是显存。24GB的显卡跑batch size 12、crop 256x512是舒服的但如果是12GB甚至8GB的卡就需要认真做减法了。我实测下来几个可调的旋钮和它们的影响如下可调参数显存影响训练效果影响batch size线性关系最直接过小会导致梯度震荡loss不稳定crop尺寸降低分辨率影响面大特征细节减少边缘视差精度下降maxdisp192降到160大视差区域会失效如果场景视差不大影响有限3D CNN通道数修改模型结构对效果影响最明显不建议动如果你的显卡只有8GB我建议先从crop尺寸入手把256x512改成192x384batch size设置4到6。再配合梯度累积模拟更大的batch虽然训练速度慢但至少loss曲线能正常下降。另外要提醒的是不要在PSMNet上轻易开启AMP混合精度。我在实验中发现3D卷积对数值精度非常敏感半精度下loss偶尔会出现NaN而且排查起来极其痛苦。最后我直接全程FP32训练换来的是稳定。4.2 loss曲线的正常形态与异常排查PSMNet的loss由三个预测头的loss加权求和因为堆叠沙漏结构会在多个阶段输出预测。原论文的加权方式是0.5 * loss1 0.7 * loss2 1.0 * loss3代码里也是这么写的。每个子loss用的是Smooth L1 Loss计算时只考虑视差真值有效大于0的像素。我第一次完整训练时loss一直在2.5附近纹丝不动过了好几个epoch都没下降。当时第一反应是学习率有问题但换成更小的学习率后仍然不动。后来一步步排查发现是数据加载时左右图顺序对调了导致模型试图从“反了”的输入里找视差根本学不到任何东西。检查完之后我把左图和右图的对应对齐loss才开始正常下降。正常的loss下降趋势大致是前10个epoch从2.x降到1.x50个epoch后降到0.6到0.8之间后面逐渐逼近0.4左右。如果你训练到100个epoch时loss还在1.5以上大概率不是训练时间不够而是数据或者代码里有问题。4.3 学习率策略和优化器配置PSMNet原论文使用的是Adam优化器初始学习率0.001训练300个epoch在第200个epoch时把学习率降到0.0001。这个策略在我的复现中效果不错没有做太多改动。有一点值得注意如果batch size因为显存限制被调小了学习率也需要相应降低。比如batch size从12降到6建议初始学习率降到0.0005左右否则梯度噪声会明显增大loss曲线会像心电图一样乱跳。我还做了一个很小的修改在main.py里加了TensorBoard的loss曲线记录这样就能远程盯着训练状态不用每次都print一条日志看一眼。这个习惯对长训练非常有用。4.4 显存溢出时的恢复操作训练中途Out of Memory确实很烦但更烦的是恢复时的状态管理。很多人直接在main.py里像python main.py --resume那样恢复训练但搞不好会遇到优化器状态和模型状态不在同一个设备上的问题。我的做法是每隔几个epoch在本地保存一份包含model.state_dict()、optimizer.state_dict()、epoch的checkpoint文件。恢复训练时先把模型和优化器都转到GPU上再加载checkpoint顺序不能反。否则你会在恢复训练的第一个step遇到莫名其妙的device mismatch错误。5. 测试与评估视差图可视化、D1指标和模型导出5.1 训练到第几个epoch用来测试训练时不需要非得等到300个epoch结束才测试。我一般会在训练过程中每20个epoch就保存一份checkpoint然后用验证集评估一次D1指标和视差图效果。PSMNet这种大模型往往在200到300个epoch之间效果就已经很好了太早的checkpoint边缘质量很差太晚的可能会在小规模数据上轻微过拟合。一个可操作的方法先训练到150个epoch左右用验证集看一眼视差图如果边缘还不够干净继续训练到200个epoch再看。论文报告的是300个epoch的结果但实际复现中180到250个epoch就能在KITTI验证集上拿到不错的D1值。5.2 从模型输出到可视化视差图的全流程测试时模型输入是一对左右图经过前向传播后得到一个视差图。需要注意的是输入图像的归一化方式要和训练时完全一致PSMNet原代码用的是ImageNet均值和标准差MEAN [0.485, 0.456, 0.406] STD [0.229, 0.224, 0.225]模型输出的视差图是浮点类型取值范围在0到maxdisp之间。直接用plt.imshow(disp)会看到一片灰蒙蒙的图像因为显示器的默认范围是0到255而视差的数值范围通常只有0到几十。为了看得舒服通常会把视差图归一化到0到255然后套一个cv2.applyColorMap变成伪彩色图。disp_vis (disp - disp.min()) / (disp.max() - disp.min() 1e-6) disp_vis (disp_vis * 255).astype(np.uint8) disp_color cv2.applyColorMap(disp_vis, cv2.COLORMAP_INFERNO)这套可视化代码不复杂但很有用。每训练一段时间把验证集上的输出图拼在一起肉眼看一下比单独观察loss数值更能发现问题。比如边缘是否走形、遮挡区域是否全是噪声、远处物体是否糊成一片这些直观信息loss曲线给不了你。5.3 KITTI官方评估指标D1-all的计算KITTI评测的主要指标是D1-all也就是所有像素中评估误差超过阈值的像素占比。阈值是误差大于3个像素或者误差大于视差值的5%两种情况满足其一就算这个像素预测错误。计算方式可以自己写一套思路是def d1_metric(disp_pred, disp_gt, mask): disp_pred disp_pred[mask] disp_gt disp_gt[mask] err torch.abs(disp_pred - disp_gt) err torch.where(err 3.0, err, torch.zeros_like(err)) err torch.where(err (torch.abs(disp_gt) * 0.05), err, torch.zeros_like(err)) d1 (err 0).float().mean() return d1这里mask是视差真值中大于0的像素。PSMNet论文在KITTI 2015验证集上的D1-all大约是2.5%左右我在自己复现的过程中能做到3%到4%原因是训练数据只有160张、超参与原论文略有出入这个差距是正常的。千万不要因为指标和论文差了零点几个百分点就怀疑代码写错了复现实验的合理范围就在这个区间附近。6. 从零到一的最短路径一套经过验证的pipeline6.1 按严重程度排序的踩坑清单坑现象严重程度解决方案左右图加载顺序对调loss不下降或下降极慢致命确认image_2为左图image_3为右图视差真值未除以256真值范围变大训练失真致命读取时强制astype(np.float32) / 256.0CUDA版本过低PyTorch无法识别GPU致命按显卡架构选择CUDA 11.x或更高老代码未兼容Py3import阶段直接报错致命全局替换xrange、修复print语法DataLoader的num_workers过大数据加载进程崩掉或死锁严重Windows下设为0Linux下设为4到8图像归一化方式与训练不一致推理效果明显变差严重统一使用ImageNet均值标准差AMP混合精度数值不稳定loss偶发NaN严重关闭AMP全程FP32显存不足时硬上大batchOOM中等减小crop尺寸或batch size加梯度累积验证集划分不一致指标和论文对不上中等沿用0-159训练/160-199验证的划分Tensor类型不注意转换设备不匹配报错低统一在数据加载完就.cuda()6.2 最短复现路径按这个顺序执行如果你不想重复我踩过的坑直接按下面这个顺序走第一步把环境按要求建好跑smoke test确认GPU链路通。第二步把老代码的Python 3兼容性问题全部改掉确保能import模型。第三步下载KITTI 2015数据集写一个读取视差真值的函数并打印验证。第四步先用一个batch的数据在前向和反向传播上跑通也就是训练1个step看loss是不是有限值且能下降。第五步正式启动训练每20个epoch保存一个checkpoint并可视化一次视差图。第六步训练到150到200个epoch后用验证集算D1指标和论文结果对比。这套pipeline里最容易让人中途放弃的是第四步因为很多bug在单step训练时就会暴露。比如数据维度不对、模型输入输出维度不匹配、loss反向传播时梯度爆炸等用一个小batch就能快速定位不用等整个训练跑完才发现。6.3 几个提升训练效率的小经验最后分享几个我实际用下来很有用的小技巧。第一个是数据加载的优化KITTI训练集只有160张图数据量不大但每次crop、归一化的计算量不小建议在dataloader里设置num_workers4Linux下配合pin_memoryTrue能让GPU利用率从60%提升到90%以上。第二个是把验证集的视差图和真值并排保存每隔一段时间看一眼不要只盯loss。很多模型问题在视差图上是一眼就能看出来的。比如边缘有横向条纹说明代价体正则化不够远处大块区域是纯色说明对低纹理区域处理失败。第三个是如果环境允许我建议直接用RLP的Yaml或者Shell脚本把完整流程记录下来把conda环境、数据路径、训练命令、恢复训练命令全写清楚这样过了一两个月再回头看也不会浪费之前积累的经验。复现老项目最忌讳的就是“跑通了但说不清怎么跑通的”把所有关键命令固化下来才能算真正的可控可复现。我现在跑PSMNet已经比较熟练了但回头看第一次踩坑的那一周最大的教训不是某个具体报错怎么解而是意识到复现老项目必须建立“主动迁移”的心态不能指望原封不动地跑通。版本兼容、数据格式、训练策略这些都要自己重新过一遍把它当成一次再工程化。希望这篇记录能帮你把这条路走得更顺一点。
返回列表