ARMA Graph Convolutions on Cora
Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: Classify papers in the Cora citation network by topic.
ARMA Graph Convolutions on Cora
Classify papers in the Cora citation network by topic. ARMA convolutions (Bianchi et al., 2019) mimic an auto-regressive moving-average filter: several parallel stacks of propagation steps whose outputs are averaged.
Same model as PyG's examples/arma.py; trained for 200 epochs instead of PyG's 400 to keep the notebook quick.
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 NormalizeFeatures
dataset = Planetoid("data/Planetoid", name="Cora", transform=NormalizeFeatures())
data = dataset[0]
print(data)
Define the model
from k3_node.layers import ARMAConv
class ARMA(keras.Model):
def __init__(self, in_channels, hidden_channels, out_channels):
super().__init__()
self.dropout = keras.layers.Dropout(0.5)
self.conv1 = ARMAConv(in_channels, hidden_channels, num_stacks=3, num_layers=2,
shared_weights=True, dropout=0.25)
self.conv2 = ARMAConv(hidden_channels, out_channels, num_stacks=3, num_layers=2,
shared_weights=True, dropout=0.25, act=None)
def call(self, data, training=False):
x = self.dropout(data.x, training=training)
x = ops.relu(self.conv1(x, data.edge_index, training=training))
x = self.dropout(x, training=training)
return self.conv2(x, data.edge_index, training=training)
model = ARMA(dataset.num_features, 16, 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.0005),
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,
)