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