Skip to content

Point cloud segmentation with Point Transformer

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 Point Transformer

Label every point of a point cloud. Point Transformer (Zhao et al., 2021) is a U-Net of vector self-attention blocks: the encoder downsamples the cloud four times, the decoder upsamples it back, interpolating features from the coarser level and adding skip connections.

Same model as PyG's examples/point_transformer_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 PointTransformerConv, fps, knn, knn_graph, knn_interpolate
from k3_node.models import MLP
from k3_node.ops.segment import segment_max

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=10, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=10)
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 TransformerBlock(keras.layers.Layer):
    def __init__(self, in_channels, out_channels):
        super().__init__()
        self.lin_in = keras.layers.Dense(in_channels, activation="relu")
        self.lin_out = keras.layers.Dense(out_channels, activation="relu")
        self.transformer = PointTransformerConv(
            in_channels, out_channels,
            pos_nn=MLP([3, 64, out_channels], norm=None, plain_last=False),
            attn_nn=MLP([out_channels, 64, out_channels], norm=None, plain_last=False),
        )

    def call(self, x, pos, edge_index):
        return self.lin_out(self.transformer(self.lin_in(x), pos, edge_index))


class TransitionDown(keras.layers.Layer):
    def __init__(self, in_channels, out_channels, ratio=0.25, k=16):
        super().__init__()
        self.k, self.ratio = k, ratio
        self.mlp = MLP([in_channels, out_channels], plain_last=False)

    def call(self, x, pos, batch, training=False):
        clusters = fps(pos, ratio=self.ratio, batch=batch)
        sub_pos, sub_batch = ops.take(pos, clusters, axis=0), ops.take(batch, clusters, axis=0)
        neighbors = knn(pos, sub_pos, k=self.k, batch_x=batch, batch_y=sub_batch)  # [cluster, point] pairs
        x = self.mlp(x, training=training)
        out = segment_max(ops.take(x, neighbors[1], axis=0), neighbors[0], num_segments=sub_pos.shape[0])
        return out, sub_pos, sub_batch



class TransitionUp(keras.layers.Layer):
    def __init__(self, in_channels, out_channels):
        super().__init__()
        self.mlp_sub = MLP([in_channels, out_channels], plain_last=False)
        self.mlp = MLP([out_channels, out_channels], plain_last=False)

    def call(self, x, x_sub, pos, pos_sub, batch, batch_sub, training=False):
        x_interpolated = knn_interpolate(self.mlp_sub(x_sub, training=training), pos_sub, pos, k=3,
                                         batch_x=batch_sub, batch_y=batch)
        return self.mlp(x, training=training) + x_interpolated


class Net(keras.Model):
    def __init__(self, in_channels, out_channels, dim_model, k=16):
        super().__init__()
        self.k = k
        self.mlp_input = MLP([in_channels, dim_model[0]], plain_last=False)
        self.transformer_input = TransformerBlock(dim_model[0], dim_model[0])
        n = len(dim_model) - 1
        self.transition_down = [TransitionDown(dim_model[i], dim_model[i + 1], k=k) for i in range(n)]
        self.transformers_down = [TransformerBlock(dim_model[i + 1], dim_model[i + 1]) for i in range(n)]
        self.transition_up = [TransitionUp(dim_model[i + 1], dim_model[i]) for i in range(n)]
        self.transformers_up = [TransformerBlock(dim_model[i], dim_model[i]) for i in range(n)]
        self.mlp_summit = MLP([dim_model[-1], dim_model[-1]], norm=None, plain_last=False)
        self.transformer_summit = TransformerBlock(dim_model[-1], dim_model[-1])
        self.mlp_output = MLP([dim_model[0], 64, out_channels], norm=None)

    def call(self, data, training=False):
        x, pos, batch = self.mlp_input(data.x, training=training), data.pos, data.batch
        x = self.transformer_input(x, pos, knn_graph(pos, k=self.k, batch=batch))
        levels = [(x, pos, batch)]
        for down, transformer in zip(self.transition_down, self.transformers_down):  # encoder
            x, pos, batch = down(x, pos, batch, training=training)
            x = transformer(x, pos, knn_graph(pos, k=self.k, batch=batch))
            levels.append((x, pos, batch))
        x = self.mlp_summit(x, training=training)
        x = self.transformer_summit(x, pos, knn_graph(pos, k=self.k, batch=batch))
        for i in range(len(self.transition_up) - 1, -1, -1):  # decoder, with skip connections
            x_skip, pos_skip, batch_skip = levels[i]
            _, pos_sub, batch_sub = levels[i + 1]
            x = self.transition_up[i](x_skip, x, pos_skip, pos_sub, batch_skip, batch_sub, training=training)
            x = self.transformers_up[i](x, pos_skip, knn_graph(pos_skip, k=self.k, batch=batch_skip))
        return self.mlp_output(x, training=training)


model = Net(3, train_dataset.num_classes, dim_model=[32, 64, 128, 256, 512], k=16)

Train

Every point is classified; besides the accuracy, the mean intersection-over-union (IoU) of the classes is reported, as in PyG. The learning rate is halved every 20 epochs.

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=99, callbacks=[keras.callbacks.LearningRateScheduler(lambda epoch: 0.001 * 0.5 ** (epoch // 20))],
          verbose=2)

Evaluate

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