Skip to content

Image classification with SplineConv and Graclus pooling

Author: K3-Node Team
Backend: Multi-Backend
Dataset: MNISTSuperpixels
Description: Classify handwritten digits represented as graphs.

View in Colab   GitHub source


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"

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

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)

Evaluate

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