Skip to content

Point cloud classification with Point Transformer

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

Classify 3D shapes from point clouds. Point Transformer (Zhao et al., 2021) applies vector self-attention among the 16 nearest neighbors of each point, with learned encodings of their relative positions. "Transition down" blocks keep a quarter of the points (farthest point sampling) and max-pool the features of their neighbors.

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

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

Sampling the points of each layer (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 Net(keras.Model):
    def __init__(self, out_channels, dim_model, k=16):
        super().__init__()
        self.k = k
        self.mlp_input = MLP([1, dim_model[0]], plain_last=False)
        self.transformer_input = TransformerBlock(dim_model[0], dim_model[0])
        self.transition_down = [TransitionDown(dim_model[i], dim_model[i + 1], k=k) for i in range(len(dim_model) - 1)]
        self.transformers_down = [TransformerBlock(d, d) for d in dim_model[1:]]
        self.mlp_output = MLP([dim_model[-1], 64, out_channels], norm=None)

    def call(self, data, training=False):
        pos, batch = data.pos, data.batch
        x = self.mlp_input(ops.ones((pos.shape[0], 1)), training=training)  # the points have no features
        x = self.transformer_input(x, pos, knn_graph(pos, k=self.k, batch=batch))
        for down, transformer in zip(self.transition_down, self.transformers_down):
            x, pos, batch = down(x, pos, batch, training=training)
            x = transformer(x, pos, knn_graph(pos, k=self.k, batch=batch))
        return self.mlp_output(global_mean_pool(x, batch, data.num_graphs), training=training)


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

Train

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"],
    run_eagerly=True,
)
model.fit(train_loader, validation_data=test_loader, epochs=200, callbacks=[keras.callbacks.LearningRateScheduler(lambda epoch: 0.001 * 0.5 ** (epoch // 20))],
          verbose=2)

Evaluate

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