Skip to content

Point cloud segmentation with PointNet++

Author: K3-Node Team
Backend: Multi-Backend
Dataset: ShapeScenes
Description: Label every point of a point cloud.

View in Colab   GitHub source


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"

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

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)

Evaluate

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