Point cloud classification with PointNet++
Author: K3-Node Team
Backend: Multi-Backend
Dataset: GeometricShapes
Description: Classify 3D shapes from point clouds.
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"
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)