Skip to content

Point cloud segmentation with RandLA-Net

Author: K3-Node Team
Backend: Multi-Backend
Dataset: ShapeScenes
Description: Label every point of a point cloud.

View in Colab   GitHub source


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"

!pip install k3-node[examples]
import os
os.environ["KERAS_BACKEND"] = "tensorflow"

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)

Evaluate

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