Image classification with SplineConv and voxel-grid pooling
Author: K3-Node Team
Backend: Multi-Backend
Dataset: MNISTSuperpixels
Description: Classify handwritten digits represented as graphs.
Image classification with SplineConv and voxel-grid pooling
Classify handwritten digits represented as graphs. Three SplineConv layers are interleaved with
voxel_grid pooling, which lays a regular grid over the image and merges all nodes inside the same
cell. The last grid has 2x2 cells, giving every image a fixed-size vector of 4 x 64 features.
Same model as PyG's examples/mnist_voxel_grid.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, max_pool, max_pool_x, voxel_grid
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
The grid cells grow from 5 to 7 to 14 pixels (images span 0-28). max_pool merges the nodes of each cell and re-applies transform to the coarser graph.
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.conv3 = SplineConv(64, 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 = voxel_grid(data.pos, batch=data.batch, size=5, start=0, end=28)
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 = voxel_grid(data.pos, batch=data.batch, size=7, start=0, end=28)
data.x, data.edge_attr = x, None
data = max_pool(cluster, data, transform=transform)
x = ops.elu(self.conv3(data.x, data.edge_index, data.edge_attr))
cluster = voxel_grid(data.pos, batch=data.batch, size=14, start=0, end=27.99)
x, _ = max_pool_x(cluster, x, data.batch, batch_size=num_graphs, size=4) # 4 cells per image
x = ops.reshape(x, (num_graphs, 4 * 64))
x = self.dropout(self.fc1(x), training=training)
return self.fc2(x)
model = Net(train_dataset.num_features, train_dataset.num_classes)
Train
Voxel pooling decides the graph sizes on the fly, so the model runs eagerly (run_eagerly=True). The learning rate drops by 10x after 5 and after 15 epochs, as in PyG.
steps = len(train_loader)
learning_rate = keras.optimizers.schedules.PiecewiseConstantDecay(
boundaries=[5 * steps, 15 * 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=20, verbose=2)