Skip to content

Point cloud classification with PointNet++

Author: K3-Node Team
Backend: Multi-Backend
Dataset: GeometricShapes
Description: Classify 3D shapes from point clouds.

View in Colab   GitHub source


Point cloud classification with PointNet++

Classify 3D shapes from point clouds. PointNet++ (Qi et al., 2017) repeatedly picks well-spread sample points (farthest point sampling, fps), groups the points within a radius around each of them, and summarizes each group with a small PointNet (PointNetConv).

Same model as PyG's examples/pointnet2_classification.py; on GeometricShapes instead of ModelNet, to keep the example small.

Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"

!pip install k3-node[examples]
import os
os.environ["KERAS_BACKEND"] = "tensorflow"

Load the data

PyG's example uses ModelNet (thousands of CAD models). GeometricShapes is a tiny stand-in: 40 kinds of shapes (cubes, spheres, pyramids, ...), one training and one test mesh each. NormalizeScale centers and scales every mesh once; SamplePoints samples 1,024 points from its surface every time it is loaded, so the model sees a new point cloud each epoch.

import keras
from keras import ops
from k3_node.datasets import GeometricShapes
from k3_node.loader import DataLoader
from k3_node.transforms import NormalizeScale, SamplePoints
from k3_node.layers import PointNetConv, fps, global_max_pool, radius
from k3_node.models import MLP

pre_transform, transform = NormalizeScale(), SamplePoints(1024)
train_dataset = GeometricShapes("data/GeometricShapes", True, transform, pre_transform)
test_dataset = GeometricShapes("data/GeometricShapes", False, transform, pre_transform)
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=32)
print(train_dataset[0])

Define the model

Each set abstraction module keeps a fraction (ratio) of the points and summarizes up to 64 neighbors within radius r of each. Sampling the points of each layer (fps, radius, knn) happens on the host with data-dependent sizes, so the model runs eagerly.

class SAModule(keras.layers.Layer):
    def __init__(self, ratio, r, nn):
        super().__init__()
        self.ratio, self.r = ratio, r
        self.conv = PointNetConv(nn, add_self_loops=False)

    def call(self, x, pos, batch):
        idx = fps(pos, batch, ratio=self.ratio)
        pos_dst, batch_dst = ops.take(pos, idx, axis=0), ops.take(batch, idx, axis=0)
        row, col = ops.unstack(radius(pos, pos_dst, self.r, batch, batch_dst, max_num_neighbors=64))
        edge_index = ops.stack([col, row], axis=0)
        x_dst = None if x is None else ops.take(x, idx, axis=0)
        x = self.conv((x, x_dst), (pos, pos_dst), edge_index)
        return x, pos_dst, batch_dst


class Net(keras.Model):
    def __init__(self, num_classes):
        super().__init__()
        self.sa1 = SAModule(0.5, 0.2, MLP([3, 64, 64, 128]))
        self.sa2 = SAModule(0.25, 0.4, MLP([128 + 3, 128, 128, 256]))
        self.sa3 = MLP([256 + 3, 256, 512, 1024])  # global set abstraction
        self.mlp = MLP([1024, 512, 256, num_classes], dropout=0.5, norm=None)

    def call(self, data, training=False):
        x, pos, batch = self.sa1(None, data.pos, data.batch)
        x, pos, batch = self.sa2(x, pos, batch)
        x = self.sa3(ops.concatenate([x, pos], axis=1), training=training)
        x = global_max_pool(x, batch, data.num_graphs)
        return self.mlp(x, training=training)


model = Net(train_dataset.num_classes)

Train

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=0.001),
    loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
    metrics=["accuracy"],
    run_eagerly=True,
)
model.fit(train_loader, validation_data=test_loader, epochs=200, verbose=2)

Evaluate

loss, accuracy = model.evaluate(test_loader, verbose=0)
print(f"Test accuracy: {accuracy:.4f}")