Skip to content

Shape correspondence with SplineCNN

Author: K3-Node Team
Backend: Multi-Backend
Dataset: MeshCorrespondence
Description: Find which vertex of a reference mesh every vertex of a deformed mesh corresponds to, i.e. classify every vertex into one of 507 classes.

View in Colab   GitHub source


Shape correspondence with SplineCNN

Find which vertex of a reference mesh every vertex of a deformed mesh corresponds to, i.e. classify every vertex into one of 507 classes. PyG's example uses the FAUST human scans; here the meshes are randomly deformed copies of a monkey head. Six SplineConv layers (Fey et al., 2018) learn from the relative positions of neighboring vertices.

Same model as PyG's examples/faust.py; on deformed monkey meshes (MeshCorrespondence) instead of FAUST, which needs a manual download.

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

FaceToEdge connects the vertices of every mesh triangle; all vertices get the same constant feature, so the model must rely on geometry. Cartesian stores the relative position of the two ends of every edge in edge_attr.

import keras
from keras import ops
from k3_node.datasets import MeshCorrespondence
from k3_node.layers import SplineConv
from k3_node.loader import DataLoader
from k3_node.transforms import Cartesian, Compose, Constant, FaceToEdge

pre_transform = Compose([FaceToEdge(), Constant(value=1)])
train_dataset = MeshCorrespondence("data/GeometricShapes", True, transform=Cartesian(), pre_transform=pre_transform)
test_dataset = MeshCorrespondence("data/GeometricShapes", False, transform=Cartesian(), pre_transform=pre_transform)
train_loader = DataLoader(train_dataset, batch_size=1, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=1)
num_vertices = train_dataset[0].num_nodes
print(train_dataset[0])

Define the model

class Net(keras.Model):
    def __init__(self, num_vertices):
        super().__init__()
        sizes = [1, 32, 64, 64, 64, 64, 64]
        self.convs = [SplineConv(a, b, dim=3, kernel_size=5, aggr="add") for a, b in zip(sizes[:-1], sizes[1:])]
        self.lin1 = keras.layers.Dense(256, activation="elu")
        self.dropout = keras.layers.Dropout(0.5)
        self.lin2 = keras.layers.Dense(num_vertices)

    def call(self, data, training=False):
        x = data.x
        for conv in self.convs:
            x = ops.elu(conv(x, data.edge_index, data.edge_attr))
        return self.lin2(self.dropout(self.lin1(x), training=training))


model = Net(num_vertices)

Train

The learning rate drops from 0.01 to 0.001 after 60 epochs, as in PyG.

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=0.01),
    loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
    metrics=["accuracy"],
)
schedule = keras.callbacks.LearningRateScheduler(lambda epoch: 0.01 if epoch < 60 else 0.001)
model.fit(train_loader, validation_data=test_loader, epochs=100, callbacks=[schedule], verbose=2)

Evaluate

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