Spline Convolutions on Cora
Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: Classify papers in the Cora citation network by topic with spline-based convolutions (Fey et al., 2018).
Spline Convolutions on Cora
Classify papers in the Cora citation network by topic with spline-based convolutions
(Fey et al., 2018). Each edge gets a pseudo-coordinate (here the
normalized in-degree of its target node, from TargetIndegree), and the layer learns a continuous
filter over it. The nodes are re-split randomly into 500 validation and 500 test nodes.
Same model as PyG's examples/cora.py.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
data holds one graph: node features x, edges edge_index, labels y, and masks marking the training, validation and test nodes.
import keras
from keras import ops
from k3_node.datasets import Planetoid
from k3_node.loader import FullGraphDataset
from k3_node.transforms import Compose, RandomNodeSplit, TargetIndegree
dataset = Planetoid("data/Planetoid", name="Cora", transform=Compose([RandomNodeSplit(num_val=500, num_test=500), TargetIndegree()]))
data = dataset[0]
print(data)
Define the model
from k3_node.layers import SplineConv
class SplineGNN(keras.Model):
def __init__(self, in_channels, out_channels):
super().__init__()
self.dropout = keras.layers.Dropout(0.5)
self.conv1 = SplineConv(in_channels, 16, dim=1, kernel_size=2)
self.conv2 = SplineConv(16, out_channels, dim=1, kernel_size=2)
def call(self, data, training=False):
x = self.dropout(data.x, training=training)
x = ops.elu(self.conv1(x, data.edge_index, data.edge_attr))
x = self.dropout(x, training=training)
return self.conv2(x, data.edge_index, data.edge_attr)
model = SplineGNN(dataset.num_features, dataset.num_classes)
Train
FullGraphDataset feeds the whole graph to Keras; mask selects which nodes count in the loss and the accuracy.
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.01, weight_decay=0.005),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
weighted_metrics=["accuracy"],
)
model.fit(
FullGraphDataset(data, mask="train_mask"),
validation_data=FullGraphDataset(data, mask="val_mask"),
epochs=200,
verbose=2,
)