PyTorch 点云分割:PointNet 到 Point Transformer 的模型选型

三种任务类型

任务输入输出典型场景
语义分割点云每个点的类别地面/树木/汽车/行人
实例分割点云每个点的实例 ID区分两辆不同车辆
部件分割单个物体点云零件归属椅子的腿/靠背/坐垫

主流模型

PointNet(最经典)

输入 N×3,直接用 MLP 逐点提取特征,全局池化后再输出分割结果:

import torch
import torch.nn as nn

class PointNet(nn.Module):
    def __init__(self, num_classes):
        super().__init__()
        self.mlp = nn.Sequential(
            nn.Linear(3, 64),
            nn.ReLU(),
            nn.Linear(64, 128),
            nn.ReLU(),
            nn.Linear(128, 256)
        )
        self.cls = nn.Linear(256, num_classes)

    def forward(self, x):
        # x: [B, N, 3]
        feature = self.mlp(x)          # [B, N, 256]
        out = self.cls(feature)        # [B, N, num_classes]
        return out

预测时取 argmax 得到每个点的标签:

pred = out.argmax(dim=-1)  # [B, N]

其他模型选型

模型核心适合场景
PointNet++局部邻域分层学习室内、激光雷达、工业检测
DGCNNKNN 建图 + EdgeConv精度要求高的分类/分割
Point TransformerTransformer AttentionSOTA,大型场景
MinkowskiNet稀疏体素 + Sparse CNN超密点云,速度快

数据格式

点云数据通常有三种格式:

xyz        # 坐标 [x, y, z]
xyzrgb     # 坐标 + 颜色 [x, y, z, r, g, b]
xyz+intensity  # 激光雷达:坐标 + 反射强度

读入后转 Tensor:

import numpy as np
import torch

points = np.load("scan.npy")  # (N, 3)
x = torch.tensor(points, dtype=torch.float32)  # [N, 3]
x = x.unsqueeze(0)  # [1, N, 3]  add batch dim

Loss 函数

点云分割本质是逐点分类,直接用 CrossEntropyLoss:

criterion = nn.CrossEntropyLoss()

# pred: [B, num_classes, N]  or reshape to [B*N, num_classes]
# label: [B, N]

pred_flat = pred.view(-1, num_classes)  # [B*N, num_classes]
label_flat = label.view(-1)            # [B*N]

loss = criterion(pred_flat, label_flat)

常用公开数据集

数据集场景用途
ShapeNet Part单物体部件分割
S3DIS室内语义分割
SemanticKITTI自动驾驶激光雷达语义分割
ScanNetRGB-D 室内场景理解

工程推荐流程

点云采集

预处理(去噪、下采样、法向量估计)

PointNet++ / Point Transformer 训练

输出分割结果

后处理(聚类、过滤小区域)

入门推荐从 PointNet 开始,结构最简单,容易在本地小数据上跑通。确认流程后再换 PointNet++ 提升精度。