Point cloud classification with RandLA-Net
Author: K3-Node Team
Backend: Multi-Backend
Dataset: GeometricShapes
Description: Classify 3D shapes from point clouds.
Point cloud classification with RandLA-Net
Classify 3D shapes from point clouds. RandLA-Net (Hu et al., 2020) was designed for huge point clouds: it downsamples points randomly (fast), and compensates with a local feature aggregation module that encodes the geometry of each point's 16 nearest neighbors and pools them with attention.
Same model as PyG's examples/randlanet_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 MaxAggregation, MessagePassing, decimation_indices, knn_graph
from k3_node.layers.conv.utils import softmax
from k3_node.models import MLP
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
SharedMLP is an MLP with batch normalization and leaky ReLU after every layer. The local feature aggregation is a message passing layer: every message combines a neighbor's features with an encoding of its relative position, weighted by attention over the neighborhood. Sampling the points of each layer (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 Net(keras.Model):
def __init__(self, num_features, num_classes, decimation=4, num_neighbors=16):
super().__init__()
self.decimation = decimation
self.fc0 = keras.layers.Dense(8)
self.block1 = DilatedResidualBlock(num_neighbors, 8, 32)
self.block2 = DilatedResidualBlock(num_neighbors, 32, 128)
self.mlp1 = SharedMLP([128, 128])
self.max_agg = MaxAggregation()
self.mlp_classif = SharedMLP([128, 32], dropout=[0.5])
self.fc_classif = keras.layers.Dense(num_classes)
def call(self, data, training=False):
b1 = self.block1(self.fc0(data.pos), data.pos, data.batch) # the positions are the input features
b1, ptr1 = decimate(b1, data.ptr, self.decimation)
b2, _ = decimate(self.block2(*b1), ptr1, self.decimation)
x = self.max_agg(self.mlp1(b2[0]), b2[2], dim_size=data.num_graphs)
return self.fc_classif(self.mlp_classif(x, training=training))
model = Net(3, train_dataset.num_classes)
Train
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, verbose=2)