Point cloud classification with DGCNN
Author: K3-Node Team
Backend: Multi-Backend
Dataset: GeometricShapes
Description: Classify 3D shapes from point clouds.
Point cloud classification with DGCNN
Classify 3D shapes from point clouds. DGCNN (Wang et al., 2019) uses
dynamic graphs: in every layer, each point is connected to its 20 nearest neighbors in the current
feature space, and EdgeConv learns from the differences between neighboring points.
Same model as PyG's examples/dgcnn_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 DynamicEdgeConv, global_max_pool
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
class Net(keras.Model):
def __init__(self, out_channels, k=20, aggr="max"):
super().__init__()
self.conv1 = DynamicEdgeConv(MLP([2 * 3, 64, 64, 64]), k, aggr)
self.conv2 = DynamicEdgeConv(MLP([2 * 64, 128]), k, aggr)
self.lin1 = keras.layers.Dense(1024)
self.mlp = MLP([1024, 512, 256, out_channels], dropout=0.5, norm=None)
def call(self, data, training=False):
x1 = self.conv1(data.pos, data.batch, training=training)
x2 = self.conv2(x1, data.batch, training=training)
out = self.lin1(ops.concatenate([x1, x2], axis=1))
out = global_max_pool(out, data.batch, data.num_graphs)
return self.mlp(out, training=training)
model = Net(train_dataset.num_classes)
Train
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.001),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
metrics=["accuracy"],
)
model.fit(train_loader, validation_data=test_loader, epochs=200, verbose=2)