PyTorch Geometric 点云分类实战:从 DGCNN 到多卡训练
2026/9/5 20:54:02 网站建设 项目流程

PyTorch Geometric 点云分类实战:从 DGCNN 到多卡训练

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

给一个 1024 个点的三维点云分类,最直觉的做法是塞进 PointNet 式的全局 MLP——但纯全连接对局部几何不敏感,细节相近的类别(比如 ModelNet40 里的 chair 和 stool)很容易混。一个常用的折中:在点集上动态建 K 近邻图,让卷积逐层"长"出局部结构。PyTorch Geometric(PyG)把这个思路做成了开箱即用的算子,DynamicEdgeConvPointTransformerConv两个类就能分别搭出 DGCNN 和 Point Transformer 两条完整分类管线,示例代码都在examples/下,几百行可跑通。

这套库的定位是 PyTorch 上的图神经网络工具箱,对点云任务有三点直接可用的能力:

  • torch_geometric.datasets内置 ModelNet、MedShapeNet 等点云数据集,下载、切分、转换一行搞定
  • torch_geometric.nn提供 FPS 采样、kNN 建图、动态边卷积、点 Transformer 卷积等算子
  • 稀疏张量和 scatter 聚合是原生实现,批内点数不齐也不会炸

三步搭好数据加载

安装上只要两件事:装 PyG 本体,再装 pyg-lib(FPS、kNN 这些 C++ 算子的宿主):

pip install torch_geometric pip install pyg-lib # DynamicEdgeConv、fps 都依赖它,>=0.6.0

更细的依赖矩阵(CUDA 版本、可选扩展)看官方文档 docs/source/install/ 即可,不用逐行抄。

数据侧的套路是"预变换 + 在线变换"两段式:

pre_transform, transform = T.NormalizeScale(), T.SamplePoints(1024) # NormalizeScale 把点云中心化并缩放到单位球——跨样本尺度一致,卷积才稳定 # SamplePoints(1024) 在线采样到固定点数,DynamicEdgeConv 建图才不会 OOM train_dataset = ModelNet('./data/modelnet10', '10', True, transform, pre_transform) test_dataset = ModelNet('./data/modelnet10', '10', False, transform, pre_transform)

完整的数据集加载逻辑(含 MedShapeNet 按类别 7:3 切分)见 examples/dgcnn_classification.py#L42-L76。

DGCNN:图不是给好的,是每次 forward 现算的

DGCNN 的核心类是DynamicEdgeConv(源码在 torch_geometric/nn/conv/edge_conv.py):每层前向传播时,它都会基于当前坐标(或特征)重新做一次 kNN 建图,然后对每条边做"点 i 的特征 ⊕ (x_i − x_j) 的差向量"再 MLP。差向量是关键——它让消息只编码相对几何,天然平移不变。

示例里的完整网络只有两层动态卷积,见 examples/dgcnn_classification.py#L91-L108:

class Net(torch.nn.Module): def __init__(self, out_channels, k=20, aggr='max'): super().__init__() self.conv1 = DynamicEdgeConv(MLP([2*3, 64, 64, 64]), k, aggr) # k=20:每个点聚合 20 个近邻 self.conv2 = DynamicEdgeConv(MLP([2*64, 128]), k, aggr) self.lin1 = Linear(128 + 64, 1024) self.mlp = MLP([1024, 512, 256, out_channels], dropout=0.5, norm=None) def forward(self, data): pos, batch = data.pos, data.batch x1 = self.conv1(pos, batch) # 图由 3D 坐标现算 x2 = self.conv2(x1, batch) # 图由上一层特征现算 out = self.lin1(torch.cat([x1, x2], dim=1)) # 两层特征拼接,避免浅层信息丢失 out = global_max_pool(out, batch) # 变长点集 max 池化成单向量 return F.log_softmax(self.mlp(out), dim=1)

注意conv2吃的是特征x1而不是原始坐标——第二层图已经建在特征空间里,这是"动态"的含义。k 值从 20 起步,点云更稀疏或想省显存时调到 12~15,细节类任务(MedShapeNet)调大到 25 通常有小幅收益,但 kNN 的计算量随 k 线性涨。

Point Transformer:注意力 + 逐级降采样

另一条路线是PointTransformerConv(torch_geometric/nn/conv/point_transformer_conv.py):注意力权重由"特征差 + 相对位置嵌入"共同决定,位置信息显式进了 attention 的计算。示例 examples/point_transformer_classification.py 的骨架是"Transformer 块 + TransitionDown 降采样"交替堆叠,dim_model=[32, 64, 128, 256, 512]五级。

降采样部分是整篇最值得抄的片段,它把 FPS 和 kNN 组合成了标准套路(examples/point_transformer_classification.py#L64-L84):

id_clusters = fps(pos, ratio=0.25, batch=batch) # FPS(最远点采样):每步挑离已选点最远的那个,0.25 即保留 1/4 点,覆盖均匀 sub_batch = batch[id_clusters] id_k_neighbor = knn(pos, pos[id_clusters], k=16, batch_x=batch, batch_y=sub_batch) # kNN:给每个采样中心找 16 个原始点,聚合范围不丢 x_out = scatter(x[id_k_neighbor[1]], id_k_neighbor[0], dim=0, reduce='max') # 邻域取 max 作为中心特征,点数逐级 1024→256→64→16 return x_out, pos[id_clusters], sub_batch

这条路径比 DGCNN 多一套注意力开销,但下采样后深层的计算量反而更低;两个示例都在 ModelNet10 上跑 201 个 epoch、batch size 32、StepLR(step_size=20, gamma=0.5),直接在同一数据集上对比公平。

训练循环与数字锚点

训练循环没有花活,就是一个 nll_loss 的循环(examples/dgcnn_classification.py#L117-L147):

python examples/dgcnn_classification.py --dataset modelnet10 --batch_size 32

几个可以拿去对齐的数字:每样本固定 1024 点、每层 k=20、Adam lr=0.001、每 20 个 epoch 学习率减半、共 201 个 epoch。单卡上 DGCNN 的瓶颈几乎全在每层的 kNN 建图(C++ 扩展里跑,CPU/GPU 均可),Point Transformer 则在 attention 上多花一份——显存紧张时先把 batch 从 32 降到 16,比砍 k 更不伤精度。MedShapeNet 这类按类别切分的长尾数据,示例里用random.seed(42)保证 7:3 划分可复现,自己写切分时别漏这一步。

往多卡上扩:采样并行与分片

单卡吃完 1024 点数据集后,真正的扩展发生在超大图上:PyG 把邻居采样做成 RPC 服务,每张卡只持有本地分片,采样时跨卡取远程邻居(示意如下,图来自 docs/source/notes/ 的分布式章节):

入口是 examples/multi_gpu/distributed_sampling.py,用torchrun起多进程后,DistNeighborLoader会按分片表自动把"本地点"和"远程点"拆好,模型侧不用感知——这套机制对点云场景意味着:百万级点云可以按空间分片喂给多卡,而不必先塞进单卡显存。

选型建议与延伸

  • 类别靠局部形状区分(椅子 vs 桌子):DGCNN 够用,k 值是第一调参旋钮
  • 需要长程依赖或全局结构(细粒度、遮挡多):Point Transformer,接受注意力开销
  • 想快速复现对比:两个示例共用 ModelNet10/1024 点/201 epoch 设定,直接并排跑

延伸阅读:

  • examples/dgcnn_segmentation.py 同构网络做点级分割,global_pool 换成逐点输出即可
  • benchmark/points/ PointNet++、EdgeCNN、SplineCNN 在 ModelNet10 上的统一评测脚本
  • examples/multi_gpu/ 分布式采样与模型并行的完整可跑示例

有类似"点云上叠注意力"或"多卡切分点云"的场景,欢迎聊聊踩过的坑。

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询