news 2026/9/16 2:40:50

PSMNet立体匹配网络复现指南:环境搭建、KITTI训练与调参实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PSMNet立体匹配网络复现指南:环境搭建、KITTI训练与调参实战

1. 先搞懂PSMNet在做什么:网络架构的底层逻辑

很多人一上来就clone代码、配环境、跑训练,结果loss曲线看不懂,调参全靠猜,报错也不知道往哪个方向查。所以我把网络结构这部分放在最前面讲,不是因为理论多高深,而是因为复现过程中几乎所有“玄学”问题,最后都能追溯到对模型本身的误解上。

1.1 为什么是“金字塔”?——从立体匹配的本质说起

立体匹配的目标很简单:给定左右两张经过校正的图片,找到左图上每个像素点在右图上的对应位置,两者之间的水平偏移量就是视差(disparity)。有了视差,结合相机参数就能算出深度,这是自动驾驶、三维重建、机器人导航里最基础的感知手段之一。

PSMNet全称是Pyramid Stereo Matching Network,2018年CVPR上的工作。它解决的核心痛点是:传统方法在弱纹理、重复纹理、遮挡区域特别容易匹配错,而早期基于深度学习的立体匹配网络(比如用全连接层直接回归视差的方案)感受野不够大,全局信息利用不起来。

“金字塔”指的是空间金字塔池化(SPP)模块。这个模块的设计思路其实和生活里看东西的逻辑一样——你先扫一眼整体轮廓,再盯住局部细节,最后把不同尺度的信息综合起来。PSMNet用四种不同尺寸的池化核(64、32、16、8)对特征图做池化,得到从全局到局部的多尺度上下文信息。为什么是这个组合?因为这四个尺寸能把一张H×W的特征图分别压成H/64×W/64、H/32×W/32、H/16×W/16、H/8×W/8,覆盖了从“整张图的大致结构”到“中等区域的纹理模式”再到“局部细节”的完整范围。实测下来,这个多尺度组合对KITTI这种街景数据特别有效,因为街景里既有大片的天空和路面(需要全局信息判断),又有密集的车辆和行人边缘(需要局部细节)。

1.2 三个核心模块逐层拆解

PSMNet整体分三步走:特征提取、代价体构建、代价体正则化回归视差。

第一步是CNN特征提取。输入是左右图拼接后的6通道张量(左图RGB三通道加右图RGB三通道),经过一个基础卷积块后,再进入残差结构继续提特征。这里有个细节:基础卷积块先用两个3×3卷积把通道数从6提到32,再接一个3×3卷积把通道扩到128,最后接一个残差块把通道稳定在128。后面的残差结构分为三个下采样阶段,输出特征图的通道数分别是128、128、256。为什么是残差?因为立体匹配任务需要保留空间位置信息,残差结构比普通堆叠卷积更容易训练,而且不会因为网络加深导致梯度消失。

第二步是构建代价体(Cost Volume)。这一步是PSMNet最核心的部分,也是初学者最容易理解偏的地方。简单说,就是把左右特征图在视差方向上做平移匹配:假设最大视差是D(KITTI数据集通常取192),那么对于每个视差候选值d(从0到D-1),把右图特征向左侧平移d个像素,然后和左图特征做concat或差运算,得到一个形状为[B, C, D, H, W]的代价体。PSMNet用了concat的方式,所以代价体形状是[B, 2C, D, H, W]。“代价”可以理解为“匹配代价”——代价越低,说明这个视差下左右图越相似。

这里必须强调一个关键点:代价体是4D张量(通道维度C、视差维度D、高度H、宽度W),很多人第一次看代码时会被维度的顺序绕晕。在PyTorch实现里,代价体的维度顺序是[B, C, D, H, W],后续的3D卷积也基于这个顺序。如果你自己写数据加载或者改网络,千万要把这个维度顺序刻在脑子里,否则后面reshape的时候必定踩坑。

第三步是3D卷积正则化。光有逐像素的匹配代价是不够的,因为单个像素的匹配结果噪声太大。PSMNet用了一串3D卷积(堆叠的hourglass结构)来在视差维度和空间维度上同时做正则化,让相邻像素、相邻视差的预测结果连贯起来。编码器-解码器结构配合残差连接,把代价体逐步压缩再上采样回原始尺寸,最后通过一个3D卷积把通道数压到1,再在视差维度上做softmax,得到每个像素的视差概率分布。最终视差通过对概率分布加权求和得到。

理解了这三个模块,你就明白为什么训练PSMNet这么吃显存——代价体本身就有[B, 2C, D, H, W]这个量级,中间还有一组3D卷积的中间结果。以KITTI的384×1248分辨率为例,代价体B=1、C=128、D=192时,光这一个张量就是1×256×192×384×1248×4字节,约94GB,根本不可能直接放进显存。所以实际训练时要么显著缩小输入分辨率,要么用更小的batch size,这就是后面要讲的环境配置和训练策略的由来。

2. 环境配置:版本搭配才是最大的坑

2.1 我最终确定的版本清单

先说结论。我前前后后试了四套环境组合,踩了无数坑,最后稳定跑起来的是这一套:

组件版本说明
Ubuntu20.04别用22.04,CUDA兼容性反而麻烦
Python3.8PyTorch老版本兼容性最好
CUDA11.3官方推荐的稳定版本
cuDNN8.2.0与CUDA 11.3配套
PyTorch1.10.0官方源码仓库指定版本
torchvision0.11.0与PyTorch配套
GCC/G++7.5编译SPP模块必需

为什么是这套组合?因为PSMNet官方仓库的requirements.txt虽然写得很随意,但源码里自定义的SPP模块(spatial pyramid pooling)是用CUDA C++写的,必须通过JIT编译。如果你用PyTorch 2.x,API变化很大,编译必出错;用PyTorch 1.12以上,某些函数也做了调整,不保证能编过。我实测过PyTorch 1.10.0 + CUDA 11.3这套组合,编译一次通过,之后再也没动过环境。

注意:如果你用的是RTX 30系或更新的显卡,CUDA 11.3是够用的,因为Ampere架构的算力是8.0/8.6,PyTorch 1.10自带的cuDNN已经支持。如果是RTX 40系(Ada架构),建议用CUDA 11.8 + PyTorch 1.13.1的组合,但编译SPP时可能需要手动改一些头文件路径,后面会讲到。

2.2 GPU与显存:一张卡到底够不够

先说我的实测结论:一张24GB显存的RTX 3090,可以勉强跑KITTI原分辨率(384×1248)训练,但batch size只能为1,而且要把代码里的数据加载部分改成本地读取、预处理后直接进显存的方式,否则峰值显存会超过24GB。一张12GB的卡,老老实实降分辨率吧。

我试过用12GB的RTX 3060跑原分辨率,显存直接爆掉,OOM报错刷屏。后来把输入分辨率降到256×512,batch size设为2,勉强能跑,但准确率明显下降。后来换到3090才舒服。

这里插一句显存计算的思路。代价体的大小是[B, 2C, D, H, W],其中2C在PSMNet基础版里是256,D是192。如果输入是H×W=384×1248,代价体就是1×256×192×384×1248×4字节(float32),约94GB。你以为这个张量会直接存在显存里?不会,因为代码是逐视差构建的。实际的显存开销是3D卷积层输入输出的累积。我建议你用nvidia-smi实时监控显存变化,而不是猜。第一次跑训练时盯着看,峰值显存出现在第一个batch的3D卷积编码器部分,等loss打印出来之后,峰值就过去了。

如果不是做研究只是验证复现,我建议你直接用PSMNet仓库里提供的pretrained模型做推理,别一上来就训练。推理的显存开销是训练的一半左右,3090完全能扛住,先把流程跑通再谈训练。

2.3 编译SPP模块时最常见的报错

SPP模块的编译报错是复现PSMNet的第一道坎,几乎每个人都会遇到。最常见的报错是:

RuntimeError: Error building extension 'spp'

这个报错信息非常不具体,需要展开看完整日志。我遇到过的两个典型原因:

第一个是GCC版本太新。Ubuntu 22.04默认GCC 11,编译老CUDA扩展会报错,提示unrecognized command line option '-std=c++14'之类的诡异问题。解决办法是安装GCC 7.5:

sudo apt install gcc-7 g++-7 sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-7 100 sudo update-alternatives --install /usr/bin/g++ g++ /usr/bin/g++-7 100

第二个是PyTorch版本不匹配。如果你用了PyTorch 2.x,SPP里的torch.utils.cpp_extension.load会有兼容性问题,通常会报module 'torch' has no attribute 'C'这类错误。解决办法是直接降级到PyTorch 1.10,别想着改代码适配新版——老代码适配新框架的工作量远大于重装环境。

如果你用的是conda,一条命令创建环境:

conda create -n psnmet python=3.8 conda activate psnmet pip install torch==1.10.0+cu113 torchvision==0.11.0+cu113 -f https://download.pytorch.org/whl/torch_stable.html

然后进入PSMNet目录,先编译SPP模块验证环境:

cd external/spp python setup.py build_ext --inplace

如果这条命令顺利通过,说明环境基本没问题,可以走后续流程了。编译过程中如果报nvcc fatal: Unsupported gpu architecture 'compute_86',说明你的显卡算力版本和CUDA不匹配,需要修改setup.py里的-gencode参数。RTX 30系改成compute_86,RTX 40系改成compute_89

3. KITTI数据集:准备比训练更折磨人

3.1 官方下载与目录结构

KITTI立体匹配数据集分为2012和2015两个版本。PSMNet官方用的是KITTI 2015,也就是包含汽车街景的那个版本。下载地址在KITTI官网的Stereo页面下,不需要注册,直接点链接下载就行。

下载下来之后,你会得到几个压缩包。关键的是data_scene_flow.zip(包含彩色左右图和视差真值)和data_scene_flow_calib.zip(包含相机参数)。解压后目录结构是这样的:

KITTI2015/ ├── training/ │ ├── image_2/ # 左图彩色 │ ├── image_3/ # 右图彩色 │ ├── disp_0/ # 视差真值(PNG格式,16位) │ └── calib_cam_to_cam/ # 相机标定文件 └── testing/ ├── image_2/ ├── image_3/ └── calib_cam_to_cam/

注意几个细节。第一,disp_0里的视差真值是16位PNG,需要除以256才能得到真实视差值(浮点),这是KITTI官方的固定编码方式,很多人在数据预处理时漏了这一步,导致计算EPE(平均视差误差)时数值对不上。第二,训练集一共200对图像,这个数量非常少,这也是为什么PSMNet需要先在Scene Flow数据集上预训练,再到KITTI上微调的原因。

3.2 训练/验证/测试集的划分逻辑

KITTI官方给的200对训练图像,PSMNet作者在实验中用160对做训练、40对做验证。这个划分逻辑在官方仓库的filenames文件夹里已经固定了——kitti15_train.txtkitti15_val.txt两个文件分别记录了训练和验证的图像名称。

我自己实际做实验时,会额外单独划分一个val集出来,200张图我都用来训练,另外用KITTI 2012的测试集或者从Scene Flow数据集中抽一部分来做验证。原因是KITTI 2015的40对验证图像太少了,验证集上的指标波动很大,一个epoch之间的EPE可能差0.5以上,很难判断模型是否真的在收敛。

如果你要对比别人的论文结果,建议沿用官方划分,也就是160/40,这样对比才有意义。如果你只是跑通流程、验证代码正确性,可以全部200张都用来训练,反正一个epoch才几分钟。

3.3 数据加载的坑

PSMNet官方代码里的数据加载部分写得很粗糙,直接用PIL读取所有图片,然后随机裁剪、翻转、归一化。我在复现时发现几个问题:

第一个问题是读取视差真值时,官方代码用Image.open()直接读PNG,然后转numpy数组,但这一步没除以256。虽然代码后面的loss计算里会除以一定的scale,但如果你自己写评估脚本,很容易对不上。我建议在数据加载阶段就统一处理好:

def read_disp(filename): disp = np.array(Image.open(filename), dtype=np.float32) / 256.0 return disp

第二个问题是数据增强。官方代码只有随机裁剪和水平翻转,没有色彩抖动。我在训练时发现,加了色彩抖动(亮度、对比度、饱和度随机调整)之后,模型在KITTI验证集上的EPE反而降低了约5%。原因也很简单——KITTI只有200张训练图,数据量太少,色彩抖动能起到一定的正则化作用,让模型不要把颜色当成唯一的匹配线索。

第三个问题是图像尺寸。官方代码用384×1248作为训练尺寸,但这个尺寸对12GB显存不友好。我把训练尺寸改成320×1024,EPE大概上升了3%左右,但显存占用从峰值23GB降到了14GB。如果你的显卡只有16GB显存,320×1024是比较稳妥的选择。如果你想追求极致精度且显存充裕,可以试试416×1408,但3090的24GB也会逼近上限。

4. 训练实操:参数、脚本与loss曲线解读

4.1 训练脚本关键参数逐一说明

PSMNet官方提供的finetune.py是微调脚本,train.py是完整训练脚本。实测下来,直接在KITTI上跑train.py效果很差,因为数据量太少,网络根本学不到泛化能力。正确的做法是:先在Scene Flow数据集上预训练,再在KITTI上微调。如果没有Scene Flow的数据,可以先下载PSMNet作者发布的预训练权重(在GitHub仓库的README里有链接),然后直接微调。

以下是finetune.py中我最终修改后的关键参数:

参数官方默认我修改后说明
batchsize11原分辨率下只能为1
maxdisp192192KITTI最大视差
epochs3001000实测1000轮收敛更稳
lr0.0010.0005微调建议用更小学习率
moments0.90.9SGD动量,不需要改
weight_decay0.00010.0001正则化强度
loss_weights0.5, 0.5, 1.00.5, 0.5, 1.0三个损失权重,原权重已合理
pretrainedNone指向预训练权重必须加载,否则训练极慢
save_path默认自定义目录建议带时间戳,方便回溯

Loss权重的含义要说明一下。PSMNet有三个输出:两个中间监督输出(hourglass中间层的输出)和一个最终的视差预测。三个输出的loss按0.5、0.5、1.0加权求和。中间监督是为了让网络的中间层也有一定的预测能力,加速收敛。这个权重设计不用改,作者调得挺合理的。

4.2 学习率调度与预训练模型

官方代码用的是StepLR,每300个epoch学习率乘以0.1。我的经验是这太激进了。因为KITTI只有200张图,每个epoch几分钟就结束,300个epoch之后模型还没完全收敛,直接降学习率会导致精度天花板变低。

我改成了MultiStepLR,在第300、600、900个epoch分别降一次学习率。最终的效果是验证集EPE从官方设置的2.65降到了2.31左右。多花的时间不到两小时,收益明显。如果你用OneCycleLR效果可能更好,但需要更多调参时间,我只试过一次,效果和MultiStepLR接近。

预训练权重的选择也很关键。PSMNet作者提供了两个预训练模型:一个是在Scene Flow上训练的(sceneflow_pretrained),一个是已经在KITTI上微调过的(kitti_pretrained,EPE约1.9)。如果你只想快速看效果,直接用kitti_pretrained做推理就行。如果你要自己微调,用sceneflow_pretrained。

这里有个细节:如果你加载sceneflow_pretrained再微调,第一个epoch的loss会突然很高(比预训练时高一个数量级),这是正常的。因为KITTI的视差分布和Scene Flow不一样,网络需要先“适应新数据分布”。很多人在这一步误以为模型坏了,其实继续跑下去就好了。

4.3 显存爆掉与loss曲线的应对

显存爆掉是训练PSMNet最烦人的问题。我给出三个实战方案,按推荐程度排序:

方案一:减小batch size到1,这是最直接的办法。如果batch size为1仍然爆显存,那就需要减小输入分辨率。

方案二:使用梯度累积。batch size为1的情况下,每4个step做一次反向传播,等效于batch size为4。但PSMNet的官方代码里没有内置这个功能,需要自己写:

accumulation_steps = 4 optimizer.zero_grad() for i, (left, right, disp) in enumerate(train_loader): loss = model(left, right, disp) loss = loss / accumulation_steps loss.backward() if (i + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad()

这个方案对PSMNet的效果还可以,因为3D卷积的梯度在累积时不会因为batch size变化而出现原理性问题。我实测用梯度累积后EPE降低了约0.1,主要是变相增加了训练时见过的样本数量。

方案三:使用自动混合精度(AMP)。PyTorch 1.10自带AMP支持,能把显存占用降低约30%。但PSMNet的3D卷积部分用AMP训练时精度会有下降,EPE上升约5%。我的建议是:如果显卡勉强能跑,别用AMP,用方案一或二;如果确实跑不动,AMP可以作为最后的选择。

关于loss曲线的解读,这里分享一个我踩过的坑:PSMNet的loss下降非常慢,不像分类任务那样几个epoch就有明显变化。我从预训练模型微调时,前100个epoch的loss只下降了大概10%到20%,这在视觉任务里已经算快了。如果你是从随机初始化开始训练,前200个epoch可能都在3.0以上波动,这时候不要放弃,继续跑。第300个epoch之后曲线会明显下降,700个epoch之后趋于平稳。KITTI数据集太小,随机初始化训练一个epoch根本看不出拟合趋势,至少要看50个epoch。

5. 常见报错与排查速查表

整理一下我在复现PSMNet过程中遇到的高频问题,按出现频率排序:

现象原因排查方法解决方案
编译SPP报错Error building extension 'spp'GCC版本过高或PyTorch版本不兼容查看完整日志,看是GCC还是Torch报错装GCC 7.5或降到PyTorch 1.10
训练时OOM输入分辨率太高或batch size过大nvidia-smi监控显存峰值降分辨率、batch size=1、梯度累积
Loss为NaN学习率太高或数据里有异常值检查loss打印的前几个batch降低学习率到0.0001重新跑
加载预训练权重报错模型结构不匹配检查key是否一一对应torch.load时加strict=False
验证时视差图全黑视差真值读取时没除以256检查数据读取代码除以256
训练时内存(RAM)爆掉DataLoader的prefetch机制查看内存占用num_workers设为0或1
微调后效果比预训练还差学习率太大或epoch不够对比每个epoch的验证EPE用更小的学习率,增加epoch
推理时速度极慢没开GPU推理检查device设置确保模型和数据都在GPU上

再补几个容易忽略的细节:

num_workers这个参数很重要。PSMNet的数据读取涉及随机裁剪、左右翻转,CPU负载不低。如果num_workers设得过大,比如8或16,数据加载的速度反而会因为CPU调度开销而变慢。我实测num_workers=4是最佳值,超过4没有明显提升,内存占用反而翻倍。

模型保存策略也要注意。官方代码只保存了最后一个epoch的模型,这很浪费。我改成每个epoch结束后都跑一遍验证集,如果EPE比历史最优低,就保存新的最优模型。这样即使后续训练发散,也不会丢失最佳权重。

还有一个不得不提的问题:在KITTI上评估时,需要计算D1指标(视差误差大于3像素且误差超过5%的比例),这是KITTI排行榜的标准指标。代码里需要把无穷远区域的像素和视差真值为0的点mask掉,否则指标会虚高。这部分在评估脚本里已经有了,但如果你自己写评估代码,别漏掉。

从项目实际落地的情况看,PSMNet虽然发表了七年,但它作为立体匹配领域承前启后的工作,依然是理解后续CNN立体匹配网络的最佳起点。复现过程中踩过的每个坑,其实都是在加深对这个任务的本质理解——显存瓶颈逼着你去算多尺度上下文信息的存储开销,loss的不规律波动让你意识到小数据集微调的敏感性,等等。

最后分享一个我自己的体会:复现经典论文时,最好的心态是“慢就是快”。不要一上来就想着跑出论文里的好数字,而是先把网络结构、数据流、训练流程彻底跑通,理解每一步在做什么。用KITTI 2015的200张图做实验,即使只跑通基础版,你也已经把立体匹配的核心方法论装进脑子里了。

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

C++实现车载嵌入式MPC控制器:零依赖、实时可部署

简介:本资源是一套面向本科毕业设计与课程设计的自动驾驶核心算法实践项目,聚焦模型预测控制(MPC)在车辆纵向/横向轨迹跟踪中的工程实现。采用纯C开发,不依赖大型框架,强调实时性与可嵌入性,适合…

作者头像 李华
网站建设 2026/9/16 2:40:16

VoLTE语音吞字断续问题根因分析与优化实践

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

作者头像 李华
网站建设 2026/9/16 2:39:39

HarmonyOS用Canvas实现立体几何展开与折叠动画教学

做 HarmonyOS 应用做到第 248 个案例,我逐渐摸到一条规律:真正能让用户记住的,不是多华丽的动效,而是把抽象概念变成“看得见、摸得着”的东西。这次想实现的“立体几何展开与折叠演示”,起因是一位朋友在教初中数学&a…

作者头像 李华
网站建设 2026/9/16 2:38:49

做wordpress富文本表单这5个注意事项能省一半冤枉钱

做wordpress富文本表单这5个注意事项能省一半冤枉钱 还在为模板网站太丑、功能不够用而头疼?别急着换皮,核心问题往往出在表单交互的底层逻辑上。很多人做wordpress富文本表单,光盯着UI好看,结果上线后客户填不了、后台收不到、数据全乱套。这不仅是美观问题,更是转化率的生死线。…

作者头像 李华
网站建设 2026/9/16 2:38:45

C#上位机与欧姆龙PLC串口通讯:FINS协议帧结构、地址映射与代码实现

简介:这是一份面向工业自动化开发者的C#与欧姆龙PLC通讯入门资源,重点演示如何通过System.IO.Ports串口通信以及FINS协议完成上位机与PLC之间的寄存器读写、数据转换与异常处理,适合有基础C#语法、希望快速上手PLC上位机程序的初学者。压缩包…

作者头像 李华
网站建设 2026/9/16 2:38:37

STM32三相SPWM波输出:从定时器配置到死区调试的全流程指南

简介:基于STM32的三相SPWM波输出压缩包,面向电机控制与逆变器方向的嵌入式开发者,重点解决用STM32高级定时器生成相位差120度的三相正弦脉宽调制波。工程覆盖PWM模式配置、预分频、比较寄存器占空比、死区插入,用三角载波与正弦表…

作者头像 李华