Point cloud segmentation with PointNet++
Author: K3-Node Team
Backend: Multi-Backend
Dataset: ShapeScenes
Description: Label every point of a point cloud.
Point cloud segmentation with PointNet++
Label every point of a point cloud. PointNet++ (Qi et al., 2017) first summarizes the cloud at coarser and coarser levels (set abstraction), then propagates the features back to all points by interpolating from the 3 nearest coarse points (feature propagation), with skip connections from each level.
Same model as PyG's examples/pointnet2_segmentation.py; on ShapeScenes instead of ShapeNet (not downloadable), with the mean IoU over the scene classes.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
PyG's example segments ShapeNet airplanes into parts. ShapeScenes is a small stand-in: each scene holds three objects (cubes, spheres, cones or tori) and every point must be labeled with the kind of object it lies on (4 classes). The node features x are the surface normals. Training scenes are randomly jittered and rotated a little.
import keras
from keras import ops
from k3_node.datasets import ShapeScenes
from k3_node.loader import DataLoader
from k3_node.transforms import Compose, RandomJitter, RandomRotate
from k3_node.layers import PointNetConv, fps, global_max_pool, knn_interpolate, radius
from k3_node.models import MLP
transform = Compose([RandomJitter(0.01), RandomRotate(15, axis=0), RandomRotate(15, axis=1), RandomRotate(15, axis=2)])
train_dataset = ShapeScenes("data/GeometricShapes", train=True, transform=transform)
test_dataset = ShapeScenes("data/GeometricShapes", train=False)
train_loader = DataLoader(train_dataset, batch_size=12, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=12)
print(train_dataset[0])
Define the model
Sampling the points of each level (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 GlobalSAModule(keras.layers.Layer):
def __init__(self, nn):
super().__init__()
self.nn = nn
def call(self, x, pos, batch, num_graphs, training=False):
x = global_max_pool(self.nn(ops.concatenate([x, pos], axis=1), training=training), batch, num_graphs)
return x, ops.zeros((num_graphs, 3)), ops.arange(num_graphs, dtype=batch.dtype)
class FPModule(keras.layers.Layer):
def __init__(self, k, nn):
super().__init__()
self.k, self.nn = k, nn
def call(self, x, pos, batch, x_skip, pos_skip, batch_skip, training=False):
x = knn_interpolate(x, pos, pos_skip, batch, batch_skip, k=self.k)
if x_skip is not None:
x = ops.concatenate([x, x_skip], axis=1)
return self.nn(x, training=training), pos_skip, batch_skip
class Net(keras.Model):
def __init__(self, num_classes):
super().__init__()
self.sa1 = SAModule(0.2, 0.2, MLP([3 + 3, 64, 64, 128]))
self.sa2 = SAModule(0.25, 0.4, MLP([128 + 3, 128, 128, 256]))
self.sa3 = GlobalSAModule(MLP([256 + 3, 256, 512, 1024]))
self.fp3 = FPModule(1, MLP([1024 + 256, 256, 256]))
self.fp2 = FPModule(3, MLP([256 + 128, 256, 128]))
self.fp1 = FPModule(3, MLP([128 + 3, 128, 128, 128]))
self.mlp = MLP([128, 128, 128, num_classes], dropout=0.5, norm=None)
def call(self, data, training=False):
sa0 = (data.x, data.pos, data.batch)
sa1 = self.sa1(*sa0)
sa2 = self.sa2(*sa1)
sa3 = self.sa3(*sa2, num_graphs=data.num_graphs, training=training)
fp3 = self.fp3(*sa3, *sa2, training=training)
fp2 = self.fp2(*fp3, *sa1, training=training)
x, _, _ = self.fp1(*fp2, *sa0, training=training)
return self.mlp(x, training=training)
model = Net(train_dataset.num_classes)
Train
Every point is classified; besides the accuracy, the mean intersection-over-union (IoU) of the classes is reported, as in PyG.
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.001),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
metrics=["accuracy", keras.metrics.MeanIoU(num_classes=4, sparse_y_pred=False, name="iou")],
run_eagerly=True,
)
model.fit(train_loader, validation_data=test_loader, epochs=30, verbose=2)