Point cloud segmentation with RandLA-Net
Author: K3-Node Team
Backend: Multi-Backend
Dataset: ShapeScenes
Description: Label every point of a point cloud.
Point cloud segmentation with RandLA-Net
Label every point of a point cloud. RandLA-Net (Hu et al., 2020) downsamples the cloud four times at random and aggregates each point's neighborhood with attention, then upsamples back to all points from the nearest coarse point, with skip connections.
Same model as PyG's examples/randlanet_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 MessagePassing, decimation_indices, knn_graph, knn_interpolate
from k3_node.layers.conv.utils import softmax
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 (decimation_indices) happens on the host with data-dependent sizes, so the model runs eagerly.
def SharedMLP(channels, **kwargs):
kwargs.setdefault("act", "leaky_relu")
kwargs.setdefault("act_kwargs", {"negative_slope": 0.2})
kwargs.setdefault("norm_kwargs", {"momentum": 0.01, "eps": 1e-6})
return MLP(channels, plain_last=False, **kwargs)
class LocalFeatureAggregation(MessagePassing):
def __init__(self, channels):
super().__init__(aggr="add")
self.mlp_encoder = SharedMLP([10, channels // 2])
self.mlp_attention = SharedMLP([channels, channels], bias=False, act=None, norm=None)
self.mlp_post_attention = SharedMLP([channels, channels])
def call(self, edge_index, x, pos):
return self.mlp_post_attention(self.propagate(edge_index, x=x, pos=pos))
def message(self, x_j, pos_i, pos_j, index, size_i):
pos_diff = pos_j - pos_i
distance = ops.sqrt(ops.sum(pos_diff * pos_diff, axis=1, keepdims=True))
spatial = self.mlp_encoder(ops.concatenate([pos_i, pos_j, pos_diff, distance], axis=1))
local_features = ops.concatenate([x_j, spatial], axis=1)
scores = softmax(self.mlp_attention(local_features), index, num_nodes=size_i) # over each neighborhood
return scores * local_features
class DilatedResidualBlock(keras.layers.Layer):
def __init__(self, num_neighbors, d_in, d_out):
super().__init__()
self.num_neighbors = num_neighbors
self.mlp1 = SharedMLP([d_in, d_out // 8])
self.shortcut = SharedMLP([d_in, d_out], act=None)
self.mlp2 = SharedMLP([d_out // 2, d_out], act=None)
self.lfa1 = LocalFeatureAggregation(d_out // 4)
self.lfa2 = LocalFeatureAggregation(d_out // 2)
def call(self, x, pos, batch):
edge_index = knn_graph(pos, self.num_neighbors, batch=batch, loop=True)
x_short = self.shortcut(x)
x = self.lfa2(edge_index, self.lfa1(edge_index, self.mlp1(x), pos), pos)
return ops.leaky_relu(self.mlp2(x) + x_short, negative_slope=0.2), pos, batch
def decimate(tensors, ptr, decimation_factor): # keep a random 1/decimation_factor of each cloud's points
idx, ptr = decimation_indices(ptr, decimation_factor)
return tuple(ops.take(t, idx, axis=0) for t in tensors), ptr
class FPModule(keras.layers.Layer): # upsampling with a skip connection
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):
x = knn_interpolate(x, pos, pos_skip, batch, batch_skip, k=self.k)
return self.nn(ops.concatenate([x, x_skip], axis=1)), pos_skip, batch_skip
class Net(keras.Model):
def __init__(self, num_features, num_classes, decimation=4, num_neighbors=16):
super().__init__()
self.decimation = decimation
d = max(32, num_classes, num_features)
self.fc0 = keras.layers.Dense(d)
self.block1 = DilatedResidualBlock(num_neighbors, d, 32)
self.block2 = DilatedResidualBlock(num_neighbors, 32, 128)
self.block3 = DilatedResidualBlock(num_neighbors, 128, 256)
self.block4 = DilatedResidualBlock(num_neighbors, 256, 512)
self.mlp_summit = SharedMLP([512, 512])
self.fp4 = FPModule(1, SharedMLP([512 + 256, 256]))
self.fp3 = FPModule(1, SharedMLP([256 + 128, 128]))
self.fp2 = FPModule(1, SharedMLP([128 + 32, 32]))
self.fp1 = FPModule(1, SharedMLP([32 + 32, d]))
self.mlp_classif = SharedMLP([d, 64, 32], dropout=[0.0, 0.5])
self.fc_classif = keras.layers.Dense(num_classes)
def call(self, data, training=False):
b1 = self.block1(self.fc0(data.x), data.pos, data.batch)
b1_d, ptr1 = decimate(b1, data.ptr, self.decimation)
b2 = self.block2(*b1_d)
b2_d, ptr2 = decimate(b2, ptr1, self.decimation)
b3 = self.block3(*b2_d)
b3_d, ptr3 = decimate(b3, ptr2, self.decimation)
b4 = self.block4(*b3_d)
b4_d, _ = decimate(b4, ptr3, self.decimation)
summit = (self.mlp_summit(b4_d[0]), b4_d[1], b4_d[2])
x = self.fp1(*self.fp2(*self.fp3(*self.fp4(*summit, *b3_d), *b2_d), *b1_d), *b1)[0]
return self.fc_classif(self.mlp_classif(x, training=training))
model = Net(3, 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.01),
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)