PyG 点云处理实战:从网格到图结构,从零实现并训练 PointNet++
2026/9/12 22:16:01 网站建设 项目流程

PyG 点云处理实战:从网格到图结构,从零实现并训练 PointNet++

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

导读

本文基于 PyG(PyTorch Geometric)官方教程,系统讲解如何用图神经网络(GNN)处理点云数据:点云本身没有图结构,但借助 PyG 的SamplePointsKNNGraphRadiusGraph等变换,可以将其转换为可供全套 GNN 使用的合成图,进而在点云上完成分类与分割任务。读完本文,你将掌握点云数据集的加载与几何变换流水线、基于MessagePassing接口从零实现 PointNet++ 核心层的方法,以及完整的训练与评估流程,并能在 GeometricShapes 等数据集上直接复现约 75%–80% 的分类准确率。

背景:为什么点云需要"图"?

点云(point cloud)是一组无序的三维坐标点集合,天然不具备边(edge)结构,因此无法直接喂给基于消息传递(message passing)的图神经网络。PyG 的核心思路是:从点云合成一张图——让空间上邻近的点互相连边,从而把"点与点之间的局部几何关系"编码为图上的边。这样,GNN 的消息传递机制就可以学习有意义的局部几何结构,最终得到的点级表示可进一步用于点云分类或分割。

整个流程可以概括为三步:

  1. 加载或构造点云数据集;
  2. 通过变换把网格(mesh)均匀采样成点云;
  3. 通过 k 近邻或半径查询把点云连成图,交给 GNN 处理。

下面按这条主线展开。

3D 点云数据集

PyG 提供了多个点云数据集,包括 PCPNetDataset、S3DIS 和 ShapeNet。上手阶段,PyG 还提供了一个玩具数据集 GeometricShapes,内含立方体、球体、金字塔等 40 类几何形状。

值得注意:GeometricShapes默认保存的是网格而非点云,通过posface两个属性表示:pos存放顶点坐标,face存放三角形连接关系:

from torch_geometric.datasets import GeometricShapes dataset = GeometricShapes(root='data/GeometricShapes') print(dataset) >>> GeometricShapes(40) data = dataset[0] print(data) >>> Data(pos=[32, 3], face=[3, 30], y=[1])

从源码看,geometry.py 在process_set中会读取原始.off网格文件,并对每个网格做去中心化处理(data.pos = data.pos - data.pos.mean(dim=0, keepdim=True)),同时为每个形状打上类别标签data.y。整个数据集包含 80 个图(训练/测试各 40),平均约 148.8 个节点、859.5 条边,输入特征数为 3(即三维坐标),类别数为 40。

下图展示了数据集中第一个网格的可视化结果,可以看到它表示一个圆盘状几何体:

从网格到点云:SamplePoints 变换

由于我们要处理的是点云,第一步是把网格转换成点。PyG 提供了 SamplePoints 变换,它会按三角形面面积加权,在网格表面均匀采样固定数量的点——面片面积越大,被采样到的概率越高。

使用方法非常简单,只需把它赋给数据集的transform属性:

import torch_geometric.transforms as T dataset.transform = T.SamplePoints(num=256) data = dataset[0] print(data) >>> Data(pos=[256, 3], y=[1])

注意两点:

  • 现在示例包含256 个点,而face中保存的三角形连接关系已被移除(Data对象不再含face属性);
  • 采样是随机的:每次访问数据集时都会重新调用变换,因此每次拿到的点云都可能不同——这对训练阶段的随机增强很有用。

从 sample_points.py 的实现看,SamplePoints还提供两个可选参数:

参数类型默认值说明
numint必填采样的点数
remove_facesboolTrue若为False,保留face张量不删除
include_normalsboolFalse若为True,额外为每个采样点计算法向量并写入data.normal

其内部实现逻辑是:先对坐标归一化,用叉积计算每个三角形面的面积,得到按面积归一化的采样概率分布,再用torch.multinomial按概率抽取面片,最后用重心坐标(barycentric coordinates)在面片内部均匀插值出采样点。对应测试见 test_sample_points.py,其中验证了include_normals=True时会额外生成data.normal属性。

下图是采样后的结果——256 个点均匀分布在原网格表面上:

从点云到图:KNNGraph 与 RadiusGraph

获得点云后,下一步是构建图。因为我们要学习的是局部几何结构,所以理想情况是让空间上相近的点互连。通常有两种做法:

  • k 近邻搜索(k-NN):每个点连接距离最近的 k 个点;
  • 球查询(ball query):连接与查询点距离小于某个半径 r 的所有点。

PyG 分别通过 KNNGraph 和 RadiusGraph 两个变换实现。可以把它们与SamplePoints组合成一条完整的变换流水线:

from torch_geometric.transforms import SamplePoints, KNNGraph dataset.transform = T.Compose([SamplePoints(num=256), KNNGraph(k=6)]) data = dataset[0] print(data) >>> Data(pos=[256, 3], edge_index=[2, 1536], y=[1])

可以看到,data对象现在多了edge_index表示,共 1536 条边——正好是 256 个点每个 6 条边。下图确认了构建出的图结构符合预期:

KNNGraph 参数详解

从 knn_graph.py 源码看,KNNGraph支持以下参数:

参数类型默认值说明
kint6每个节点的邻居数
loopboolFalseTrue时图包含自环
force_undirectedboolFalseTrue时把有向边转成无向边(通过to_undirected
flowstr"source_to_target"配合消息传递的流向,为"source_to_target"时每个目标节点恰好有 k 个源节点指向它
cosineboolFalseTrue时用余弦距离代替欧氏距离找近邻
num_workersint1计算近邻使用的并行 worker 数(batch 非 None 或输入在 GPU 上时无效)

注意:KNNGraph会把data.edge_attr置为None,因为它构建的是纯几何邻接图。

RadiusGraph 参数详解

radius_graph.py 用于球查询,参数为:

参数类型默认值说明
rfloat必填连接半径(距离阈值)
loopboolFalse是否包含自环
max_num_neighborsint32每个元素最多返回的邻居数,主要用于 CUDA 张量
flowstr"source_to_target"消息传递流向
num_workersint1并行 worker 数

与 k-NN 相比,球查询的优势在于邻居数量随局部点密度自适应变化,不会强制固定 k 个邻居,这在点密度不均的场景(如真实扫描数据)中更符合物理直觉。

PointNet++ 实现

核心思想:分组 → 邻域聚合 → 下采样

PointNet++ 是点云分类与分割领域的开创性工作,它通过"分组(grouping)、邻域聚合(neighborhood aggregation)、下采样(downsampling)"三步循环,迭代式地处理点云:

  1. 分组阶段:按上文介绍的 k 近邻或球查询构建图;
  2. 邻域聚合阶段:执行一个 GNN 层,对每个点聚合其直接邻居的信息,从而在不同尺度上捕获局部上下文;
  3. 下采样阶段:实现对大小不同的点云适用的池化方案。

下图展示了 PointNet++ 的分层处理架构:

由于篇幅原因,教程中暂不实现下采样阶段;完整的下采样(最远点采样 FPS + 球查询分组)实现可参考仓库中的 examples/pointnet2_classification.py,本文末尾会简要分析。

邻域聚合:用 MessagePassing 从零实现 PointNet 层

PointNet++ 层的消息传递公式为:

$$ \mathbf{h}^{(\ell + 1)}i = \max{j \in \mathcal{N}(i)} \textrm{MLP} \left( \mathbf{h}_j^{(\ell)}, \mathbf{p}_j - \mathbf{p}_i \right) $$

其中:

  • $\mathbf{h}_i^{(\ell)} \in \mathbb{R}^d$ 表示点 $i$ 在第 $\ell$ 层的隐藏特征;
  • $\mathbf{p}_i \in \mathbb{R}^3$ 表示点 $i$ 的空间位置。

语义很清晰:对每个邻居 $j$,把"邻居特征 $\mathbf{h}_j$"与"相对位置 $\mathbf{p}_j - \mathbf{p}_i$"拼接后过 MLP 生成消息,再用max 聚合取邻域内的最大值。这里相对位置编码了局部几何结构——这正是点云 GNN 的关键先验。

PyG 的 MessagePassing 接口会自动处理消息传播,我们只需定义message函数并指定聚合方式。下面是完整的层实现:

from torch import Tensor from torch.nn import Sequential, Linear, ReLU from torch_geometric.nn import MessagePassing class PointNetLayer(MessagePassing): def __init__(self, in_channels: int, out_channels: int): # Message passing with "max" aggregation. super().__init__(aggr='max') # Initialization of the MLP: # Here, the number of input features correspond to the hidden # node dimensionality plus point dimensionality (=3). self.mlp = Sequential( Linear(in_channels + 3, out_channels), ReLU(), Linear(out_channels, out_channels), ) def forward(self, h: Tensor, pos: Tensor, edge_index: Tensor, ) -> Tensor: # Start propagating messages. return self.propagate(edge_index, h=h, pos=pos) def message(self, h_j: Tensor, pos_j: Tensor, pos_i: Tensor, ) -> Tensor: # h_j: The features of neighbors as shape [num_edges, in_channels] # pos_j: The position of neighbors as shape [num_edges, 3] # pos_i: The central node position as shape [num_edges, 3] edge_feat = torch.cat([h_j, pos_j - pos_i], dim=-1) return self.mlp(edge_feat)

逐段拆解这个实现:

  • __init__:通过super().__init__(aggr='max')指定max 聚合;随后初始化一个 MLP,其输入维度为"节点特征维度 + 3(三维坐标)",负责把"邻居特征 + 源点到目标点的相对空间关系"映射成可训练的消息;
  • forward:基于edge_index调用self.propagate(...)启动消息传播,并把消息构造所需的一切(hpos)传入;
  • message:利用 PyG 自动追加的*_j/*_i后缀,分别访问邻居节点与中心节点的信息(h_jpos_jpos_i),对每条边返回一条消息——这里即cat([h_j, pos_j - pos_i])再过 MLP。

PyG 的这一机制让实现点云 GNN 层变得非常直接:你只需关心"如何为一条边生成消息"和"如何聚合",其余传播逻辑全部交给框架。

事实上,PyG 已内置了等价的 PointNetConv 层,可直接使用(见下文"更进一步"一节)。

网络架构:两层消息传递 + 全局池化

有了PointNetLayer,就可以搭建完整的分类网络。整体架构由两层 PointNet 卷积和一个线性分类器组成:

from torch_geometric.nn import global_max_pool class PointNet(torch.nn.Module): def __init__(self): super().__init__() self.conv1 = PointNetLayer(3, 32) self.conv2 = PointNetLayer(32, 32) self.classifier = Linear(32, dataset.num_classes) def forward(self, pos: Tensor, edge_index: Tensor, batch: Tensor, ) -> Tensor: # Perform two-layers of message passing: h = self.conv1(h=pos, pos=pos, edge_index=edge_index) h = h.relu() h = self.conv2(h=h, pos=pos, edge_index=edge_index) h = h.relu() # Global Pooling: h = global_max_pool(h, batch) # [num_examples, hidden_channels] # Classifier: return self.classifier(h) model = PointNet()

打印模型可以看到所有模块都正确初始化:

print(model) >>> PointNet( ... (conv1): PointNetLayer() ... (conv2): PointNetLayer() ... (classifier): Linear(in_features=32, out_features=40, bias=True) ... )

架构要点:

  • 网络继承自torch.nn.Module,构造函数中初始化两个PointNetLayer模块和一个线性分类器
  • forward中依次应用两层图卷积,中间穿插 ReLU 非线性激活。第一层输入 3 维特征(即节点坐标),输出 32 维;第二层后每个点已经聚合了其2 跳邻域的信息,足以区分简单的局部形状;
  • 随后调用全局读出的 global_max_pool,沿节点维度对每个样本取最大值,得到[num_examples, hidden_channels]的图级表示;
  • 为了把不同节点映射回各自所属的样本,需要batch向量——它由 PyG 的 DataLoader 在 mini-batch 训练时自动创建(同 batch 内的多个图会被拼成一个大连通块,batch记录每个节点属于哪个图);
  • 最后用线性分类器把每片点云的 32 维全局特征映射到 40 个类别之一。

训练与评估

训练流程与标准 PyTorch 完全一致,只需注意数据加载与batch向量的使用。下面先构造训练/测试数据集和 DataLoader:

from torch_geometric.loader import DataLoader train_dataset = GeometricShapes(root='data/GeometricShapes', train=True) train_dataset.transform = T.Compose([SamplePoints(num=256), KNNGraph(k=6)]) test_dataset = GeometricShapes(root='data/GeometricShapes', train=False) test_dataset.transform = T.Compose([SamplePoints(num=256), KNNGraph(k=6)]) train_loader = DataLoader(train_dataset, batch_size=10, shuffle=True) test_loader = DataLoader(test_dataset, batch_size=10)

注意这里GeometricShapes通过train=True/False分别加载训练集与测试集,并各自应用相同的"采样 + 建图"变换。之后定义模型、优化器与损失函数:

model = PointNet() optimizer = torch.optim.Adam(model.parameters(), lr=0.01) criterion = torch.nn.CrossEntropyLoss()

训练循环:每个 batch 中,把data.posdata.edge_indexdata.batch送入模型得到 logits,与标签data.y计算交叉熵损失,反向传播并更新参数;损失按图数加权后取平均,得到每个 epoch 的平均损失:

def train(): model.train() total_loss = 0 for data in train_loader: optimizer.zero_grad() logits = model(data.pos, data.edge_index, data.batch) loss = criterion(logits, data.y) loss.backward() optimizer.step() total_loss += float(loss) * data.num_graphs return total_loss / len(train_loader.dataset)

测试循环:在torch.no_grad()下前向推理,取 logits 的 argmax 作为预测类别,统计正确率:

@torch.no_grad() def test(): model.eval() total_correct = 0 for data in test_loader: logits = model(data.pos, data.edge_index, data.batch) pred = logits.argmax(dim=-1) total_correct += int((pred == data.y).sum()) return total_correct / len(test_loader.dataset)

最后迭代 50 个 epoch:

for epoch in range(1, 51): loss = train() test_acc = test() print(f'Epoch: {epoch:02d}, Loss: {loss:.4f}, Test Acc: {test_acc:.4f}')

预期结果:使用上述配置,即便每个类别只有 1 个训练样本(GeometricShapes 训练集共 40 个样本、40 类),测试集准确率也能达到约75%–80%——这充分说明"几何变换 + 消息传递"这一范式在数据极少的情况下也能学到有效的局部几何特征。

更进一步:内置 PointNetConv 与完整 PointNet++ 示例

教程中的PointNetLayer是教学用实现。生产场景下,可以直接使用 PyG 内置的 PointNetConv,它支持:

  • local_nn:可选的 MLP(对应消息构造网络 $h_{\Theta}$),输入为[-, in_channels + num_dimensions],输出[-, out_channels]
  • global_nn:可选的 MLP(对应聚合后处理网络 $\gamma_{\Theta}$),输入为[-, out_channels],输出[-, final_out_channels]
  • add_self_loops:默认为True,自动为图添加自环;
  • 聚合方式默认为aggr='max',与上述公式一致。

PointNetConv的消息构造与教程实现完全对应:message中先算pos_j - pos_i,再与邻居特征x_j拼接,过local_nn。其测试见 test_point_conv.py,覆盖了稠密邻接、SparseTensor 稀疏邻接、bipartite(二部图)消息传递以及 TorchScript 编译等多种场景。

若想复现完整的 PointNet++ 层级结构(含下采样),可以参考 examples/pointnet2_classification.py:

  • 它用fps(最远点采样,Farthest Point Sampling)按比例ratio选出子集点作为下采样;
  • radius球查询在子集点周围构建局部邻域,配合PointNetConv完成分组聚合(即SAModule);
  • 三个SAModule逐层缩小点云规模、扩大感受野,最后用GlobalSAModule做全局池化读出图级特征;
  • 数据集使用 ModelNet10,配合T.NormalizeScale()(尺度归一化)与T.SamplePoints(1024)(每片点云采样 1024 点)两个变换。

小结

本文完整走通了 PyG 处理点云数据的标准范式:

阶段关键组件作用
数据准备GeometricShapes 等点云数据集提供网格/点云形式的 3D 数据
网格 → 点云SamplePoints按面面积加权均匀采样,支持法向量输出
点云 → 图KNNGraph / RadiusGraphk 近邻或球查询建边,编码局部几何关系
特征学习PointNetLayer / PointNetConv基于消息传递聚合邻域信息
图级读出global_max_pool将节点特征聚合为整图表示
完整 PointNet++pointnet2_classification.py加入 FPS 下采样与球查询分组的完整层级实现

这套"变换建图 + 消息传递"的思路并不局限于点云分类:把SamplePoints换成其他数据源、把KNNGraph换成自定义建图规则,即可将同一套 GNN 工具链推广到更广泛的几何学习任务(如点云分割、法向量估计等)。基于以上代码,你可以直接在本地运行并验证约 75%–80% 的测试集准确率,进而以此为起点深入 PointNet++ 的层级化设计。

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

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

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

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

立即咨询