ARTICLE DETAIL

资讯详情

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

PointMLP实战:PyTorch实现3D点云分类从零跑通

PointMLP实战:PyTorch实现3D点云分类从零跑通 最近后台不少读者问3D点云分类到底从哪下手网上资料要么是纯理论讲Transformer要么是论文公式堆到看不懂。这次我把Paper With Code上挺火的PointMLP拉出来用PyTorch实际跑一遍完整代码放在GitHub上整个过程5分钟就能看到训练日志。这篇文章不绕弯子直接讲清楚四件事PointMLP解决什么问题、环境怎么搭、代码怎么跑通、以及真正干活时那几个容易踩的坑。想入门点云分类的同学按这篇文章的步骤走基本不会卡住。1. 认清PointMLP它凭什么能火1.1 从PointNet到PointMLP一条清晰的技术演进线很多人第一次接触点云分类都是从PointNet开始。PointNet用共享多层感知器直接对每个点做特征提取然后通过最大池化聚合全局特征思路很直接但代价是忽略了局部结构。后来PointNet引入多尺度分组和采样把局部邻域信息抓回来了精度明显提升。不过PointNet的特征提取依然依赖人工设计的采样和分组策略遇到密度不均、噪声多的点云鲁棒性还是会打折扣。PointMLP这个工作走的是另一个方向不做复杂采样和分组干脆把局部特征提取换成堆叠的残差MLP然后通过几何仿射模块对特征做标准化。这样一来网络结构变得极其简单不需要FPN那样的复杂连接却在ModelNet40上取得了93.2%左右的准确率刷新了当时MLP类方法的记录。这个思路说白了就是别把简单问题复杂化用足够深的MLP配合理的几何约束就能把点云特征学好。对于刚入门的人来说这个模型结构比Transformer系列更容易读懂代码量也更短非常适合作为第一个跑通的点云分类项目。1.2 PointMLP在ModelNet40上的表现和硬件门槛从实际复现结果看PointMLP在ModelNet40分类任务上模型参数量大约12.7M不算大。在单张RTX 3090上训练一个epoch大概40秒左右当然这个时间跟数据加载和机器配置有关系。相比之下PointNet需要更多调参技巧才能达到接近的精度训练过程也没有PointMLP稳定。我自己在2080Ti上跑过显存占用在6GB以内说明这个模型对显卡的要求没那么苛刻普通游戏卡也能训练这点对学生党很友好。顺便说一句PointMLP的原始仓库是基于PyTorch实现的代码风格比较学术适合阅读。但是直接拿来用的话要注意数据加载部分和日志输出写得比较简单后面我会讲怎么改造成自己的训练流程。2. 动手前准备环境、数据和代码结构2.1 环境配置PyTorch与CUDA的版本搭配实战第一件事不是写代码而是把环境弄好。我建议直接用Anaconda创建虚拟环境Python版本选3.8或3.9都可以。PyTorch版本建议1.10以上2.0更好因为新版本对自动混合精度支持更完善。CUDA建议选11.7或11.8比较稳定。如果机器没有NVIDIA显卡用CPU版本PyTorch也能跑通只是训练速度会慢很多5分钟那个说法是基于有GPU的情况。创建环境的命令大概是conda create -n pointmlp python3.9 conda activate pointmlp pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install numpy scipy tqdm tensorboard注意PyTorch安装别用默认源国内网络环境经常卡住直接用官方提供的index-url会快一些。tensorboard不是必须的但是加上它方便看loss曲线我后面会用它监控训练。2.2 从GitHub上找代码怎么判断一个项目值不值得跑GitHub上关于PointMLP的仓库不少有原版也有第三方复现。我建议优先看star数较高的原版仓库因为网络上的复现版本可能改动过网络结构跑出来的结果会不一样。下载代码有两种方式一是直接在仓库页面点Code按钮下载zip压缩包二是用git clone命令拉取推荐后者方便跟着上游更新。代码下下来之后先别急着跑花两分钟看清目录结构。一般这个项目包含main.py、pointmlp.py、data_utils.py等几个核心文件。main.py负责训练和测试逻辑pointmlp.py定义网络模型data_utils.py里是数据加载和增强的逻辑。搞清楚这三个文件后面改参数就有方向了。2.3 ModelNet40数据集手动下载与目录约定PointMLP训练最常用的数据集是ModelNet40包含40个类别、9843个训练样本和2468个测试样本。原作者提供的脚本可以直接下载但是国内从Princeton服务器下载经常很慢这时候我建议用第三方镜像或者其他开源渠道把modelnet40_ply_hdf5_2048.zip下载下来解压到data文件夹下。数据集目录结构要跟代码里的读取逻辑匹配。以主流的HDF5版本为例解压后应该有ply_data_train0.h5到ply_data_train4.h5以及ply_data_test0.h5到ply_data_test1.h5这些文件。如果你的路径和代码不一致训练脚本会直接报文件找不到这是新手最容易踩的坑。3. 五分钟跑通核心实操流程3.1 训练参数选择与启动命令环境准备好、数据集放对位置之后训练就特别简单。官方默认用ModelNet40batch size设为32epoch设300初始学习率0.1用余弦退火策略下降。虽然训练脚本里默认300个epoch但实际跑起来不需要等那么久前50个epoch就已经能看到收敛趋势。如果你只是验证代码能否跑通可以加一个--epoch 5参数先把流程走一遍确认没问题再开完整训练。一个比较实用的启动命令示例python main.py --model pointmlp --epoch 300 --batch_size 32 --lr 0.1如果显存不够把batch size降到16或8即可。PointMLP对batch size的敏感度不算高小batch也能训练只是收敛速度会变慢精度会稍微受影响。3.2 测试评估指标怎么看训练完成后脚本会在测试集上输出整体准确率Overall Accuracy和平均准确率Average Accuracy。这两个指标有区别整体准确率是全部样本里预测正确的比例平均准确率是每个类别准确率加和再取平均。ModelNet40类别基本均衡所以两个指标差别不会特别大。如果你的应用场景类别不均衡就得更看重平均准确率。我自己跑出来的结果在未加vote的测试模式下整体准确率有92.8%左右跟论文里的93.2%已经很接近了。差异主要来自数据增强、随机种子和训练时长不用太纠结这0.5%的差距。3.3 推理真正对单个点云做预测很多人训练完就结束了忘了模型最终用在外界数据上的时候输入和训练时的预处理必须一致。官方代码的test函数是在完整测试集上做评估如果你想对单个点云文件做推理需要自己写一小段脚本。大致的思路是加载预训练权重读取点云文件把点云归一化到单位球或者和训练时相同的尺度然后转成tensor过一层模型用argmax取类别索引。这里有个细节要注意训练时输入的点是2048个如果你的测试点云超过2048个要先下采样如果少于2048个建议重复采样补齐。否则直接输入模型会报维度错误因为你改了输入维度之后网络里的特征维度也跟着对不上。4. 深入原理五步拆解PointMLP核心模块4.1 局部几何特征提取近邻聚合的初衷PointMLP虽然没有用PointNet那种复杂的多尺度分组但它依然依赖局部邻域信息。具体做法是在每个点周围找K个最近邻点然后把这些邻居点的坐标和中心点的坐标拼接起来再经过残差MLP提取特征。找最近邻的算法用到了KDTree或者Ball Query这两种方式各有优缺点KDTree适合点数固定的场景Ball Query对密度变化更鲁棒。PointMLP的实现中K可以理解为一个超参数通常设为32左右。为什么要做局部特征提取因为点云本身是无序的如果我们把所有点直接拉平输入全连接层网络完全不知道哪些点靠得近特征表达能力会大打折扣。局部聚合相当于给网络提供了一种先验空间上相近的点特征应该一起被抽象。4.2 残差前馈网络为什么堆层不会退化PointMLP的特点就是它的基本模块是残差前馈网络每个模块由两个1x1卷积层等价于全连接层、BatchNorm和ReLU组成模块输出加上输入形成shortcut连接。这种残差结构在2D图像里已经被ResNet验证过了放到点云上一样有效。如果没有shortcut层数越深越容易出现梯度消失训练误差反而变大加了shortcut之后梯度能直接传到前面的层网络加深也就不可怕了。实现时有一个小技巧第一个1x1卷积会把通道数升到一个较大值比如128第二个卷积会降回原始通道数。这样既保证了特征的表达能力又不会让模块输出维度和输入对不上。BatchNorm放在卷积之后、激活之前这在训练时能让数据分布更稳定但推理时要注意把BN层设置成eval模式否则结果会有细微差别。4.3 几何仿射模块对点云做标准化几何仿射模块Geometric Affine Module是PointMLP里比较有辨识度的设计。普通PointNet只对每个点做MLP和池化没有考虑到局部区域的几何特征可能存在尺度不一致。Geometric Affine的做法是先计算局部区域特征的均值和标准差然后对特征做类似Instance Normalization的处理用可学习的缩放因子和偏置项对归一化后的特征做仿射变换。这个模块的价值在于提高了对点云尺度变化的鲁棒性。实际场景中同一个物体扫描仪距离远近不同点云的尺度可能差好几倍。如果不做这类标准化网络只能见过训练时的尺度测试时碰到不同尺度的物体精度会掉得很快。4.4 分组与采样策略怎么控制计算量由于PointMLP没有像PointNet那样在每一层都用FPS采样减少点数它的计算量主要集中在前面几层的局部特征提取上。实际代码中PointMLP会在最开始把输入点云做一次分组比如每32个点一组组数等于点的数量。这样一来当点的数量从2048降到后面层时特征数量其实已经压缩过一轮整体计算量可控。如果觉得显存紧张可以适当减少batch size或者减少每组的邻居数量K不建议减少输入点数。因为输入点数直接关系到局部特征的质量2048个点是精度和速度之间比较平衡的选择。4.5 分类头与损失函数最后怎么得到类别网络的分类头很简单经过若干层特征提取后对每个点的特征做全局最大池化得到一个1024维全局特征向量然后接两层全连接分类器输出维度是类别数。损失函数用的是交叉熵损失配合标准的SGD优化器或者AdamW都可以。原代码默认用SGD并用余弦退火调整学习率这个组合在分类任务上表现不错。我个人测试时发现AdamW在训练前期收敛更快但最终精度和SGD差不多。如果你是跑完整训练建议沿用原config的SGD设置复现率最高如果你只是做快速实验AdamW会更友好两三分钟就能看到loss下降的趋势。5. 踩坑实录与性能优化5.1 训练不收敛或精度偏低如果你训练10个epoch后损失还在1.8以上大概率是学习率设得太大了。PointMLP初始学习率0.1配合batch size 32和余弦退火是合理的但如果你用了AdamW却不小心沿用0.1这个学习率loss会直接爆炸。不同优化器对学习率的敏感度完全不同这点一定要记得同步调整。另外数据归一化也很关键。ModelNet40的坐标通常落在[-1, 1]之间但如果数据集的坐标系定义不同比如毫米级的点云网络训练会非常困难。在训练脚本里查看数据加载部分确保输入模型之前做了除以外围半径的归一化否则一切后续操作都是白费。5.2 显存不足与OOM显存溢出是跑点云模型最常见的问题。解决办法不外乎三条第一降低batch size从32降到16通常就能解决大部分OOM第二开启自动混合精度AMPPyTorch的torch.cuda.amp可以大幅减少显存占用训练速度还有提升第三实在不行就减少输入点云点数但前面也说了不建议降到1024以下会影响精度。实际操作中我习惯先把batch size调小把代码跑通确认逻辑无误后再逐步增大batch size直到接近显存上限。这样的好处是能快速定位到底是代码问题还是资源问题。5.3 点云预处理和数据增强的影响点云的分类性能很大程度上依赖数据增强策略。训练时常见的增强手段包括随机旋转、随机平移、随机缩放等。旋转增强对ModelNet40的提升很明显因为这个数据集里大部分物体有明确朝向如果你把所有物体都固定朝向模型会学到过于简单的“位置记忆”。在训练代码中data_utils.py里已经实现了一套基础增强流程建议不要随意关掉。测试和推理阶段则要关掉所有随机增强保证结果可复现。很多新手把训练用的增强函数直接套用在测试集上导致精度忽高忽低还以为是模型问题。5.4 复现论文精度的小技巧想完全复现论文里的93.2%准确率除了参数保持一致外有几个小地方容易被忽略。第一随机种子很多学术代码里会设置固定的random seed如果你改过数据加载的顺序精度会有一点波动但不会显著影响结论。第二投票测试voting机制官方仓库中test阶段可以对同一物体做多次变换取平均概率作为最终预测这种方式通常能提升0.3%~0.5%的准确率代价是推理时间增加。第三训练时用余弦退火不要提前中断尽量跑满epoch数后期loss下降虽然平缓但对精度的贡献还是有的。我自己的体会是PointMLP值得作为入门点云分类的第一个完整项目因为它代码短、依赖少、训练快而且能把分类的整套流程走通。跑通之后你再去读PointNet或者PointTransformer就有了对比的基准对网络结构的理解也会更具体。最后分享一个小技巧遇到实验效果跟论文对不上时先检查数据加载和预处理是不是和论文一致大多数情况问题不在模型结构本身。这篇文章就到这里建议你直接去GitHub把代码拉下来自己动手跑一遍遇到问题多打印几个中间变量看看比盯着公式猜要快得多。
返回列表