Point cloud segmentation with Point Transformer
Author: K3-Node Team
Backend: Multi-Backend
Dataset: ShapeScenes
Description: Label every point of a point cloud.
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"
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)