Skip to content

Point cloud classification with RandLA-Net

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 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"

!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 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)

Evaluate

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