Skip to content

Image classification with NNConv 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 NNConv and Graclus pooling

Classify handwritten digits represented as graphs. NNConv (Gilmer et al., 2017) uses a small neural network to turn each edge's relative position into a weight matrix for the message along that edge; graclus clustering coarsens the graph after each convolution.

Same model as PyG's examples/mnist_nn_conv.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 NNConv, 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])


def edge_network(in_channels, out_channels):  # maps an edge's 2D relative position to a weight matrix
    return keras.Sequential([keras.layers.Dense(25, activation="relu"), keras.layers.Dense(in_channels * out_channels)])


class Net(keras.Model):
    def __init__(self, in_channels, num_classes):
        super().__init__()
        self.conv1 = NNConv(in_channels, 32, edge_network(in_channels, 32), aggr="mean")
        self.conv2 = NNConv(32, 64, edge_network(32, 64), aggr="mean")
        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}")