Image classification with NNConv and Graclus pooling
Author: K3-Node Team
Backend: Multi-Backend
Dataset: MNISTSuperpixels
Description: Classify handwritten digits represented as graphs.
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"
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)