Image classification with SplineConv and Graclus pooling
Author: K3-Node Team
Backend: Multi-Backend
Dataset: MNISTSuperpixels
Description: Classify handwritten digits represented as graphs.
Image classification with SplineConv and Graclus pooling
Classify handwritten digits represented as graphs. Two SplineConv layers
(Fey et al., 2018) learn filters over the relative positions of
neighboring nodes, and graclus clustering coarsens the graph after each of them.
Same model as PyG's examples/mnist_graclus.py; on scikit-learn's small Digits images instead of MNIST superpixels.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
PyG uses MNIST superpixels; here each image of scikit-learn's small Digits dataset becomes a graph: every non-blank pixel is a node with its intensity as feature and its position pos, linked to its neighboring pixels. 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 Digits
from k3_node.loader import DataLoader
from k3_node.transforms import Cartesian
from k3_node.layers import SplineConv, global_mean_pool, graclus, max_pool, max_pool_x
from k3_node.utils import normalized_cut
transform = Cartesian(cat=False)
train_dataset = Digits(train=True, transform=transform)
test_dataset = Digits(train=False, transform=transform)
train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=64)
print(train_dataset[0])
Define the model
After each convolution, graclus pairs up neighboring nodes (preferring pairs with a small normalized cut of their distance) and max_pool merges every pair into one node, halving the graph. max_pool also re-applies transform, so the coarser graph gets fresh edge features.
def normalized_cut_2d(edge_index, pos):
distance = ops.norm(ops.take(pos, edge_index[0], axis=0) - ops.take(pos, edge_index[1], axis=0), axis=1)
return normalized_cut(edge_index, distance, num_nodes=pos.shape[0])
class Net(keras.Model):
def __init__(self, in_channels, num_classes):
super().__init__()
self.conv1 = SplineConv(in_channels, 32, dim=2, kernel_size=5)
self.conv2 = SplineConv(32, 64, dim=2, kernel_size=5)
self.fc1 = keras.layers.Dense(128, activation="elu")
self.dropout = keras.layers.Dropout(0.5)
self.fc2 = keras.layers.Dense(num_classes)
def call(self, data, training=False):
num_graphs = data.num_graphs
x = ops.elu(self.conv1(data.x, data.edge_index, data.edge_attr))
cluster = graclus(data.edge_index, normalized_cut_2d(data.edge_index, data.pos), x.shape[0])
data = max_pool(cluster, data._replace(x=x, edge_attr=None), transform=transform)
x = ops.elu(self.conv2(data.x, data.edge_index, data.edge_attr))
cluster = graclus(data.edge_index, normalized_cut_2d(data.edge_index, data.pos), x.shape[0])
x, batch = max_pool_x(cluster, x, data.batch)
x = global_mean_pool(x, batch, num_graphs)
x = self.dropout(self.fc1(x), training=training)
return self.fc2(x)
model = Net(train_dataset.num_features, train_dataset.num_classes)
Train
graclus and pooling decide the graph sizes on the fly, so the model runs eagerly (run_eagerly=True). The learning rate drops by 10x after 15 and after 25 epochs, as in PyG.
steps = len(train_loader)
learning_rate = keras.optimizers.schedules.PiecewiseConstantDecay(
boundaries=[15 * steps, 25 * steps], values=[0.01, 0.001, 0.0001])
model.compile(
optimizer=keras.optimizers.Adam(learning_rate),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
metrics=["accuracy"],
run_eagerly=True,
)
model.fit(train_loader, validation_data=test_loader, epochs=30, verbose=2)