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