A simple GCN baseline for node property prediction
Author: K3-Node Team
Backend: Multi-Backend
Dataset: Photo (Amazon)
Description: PyG's example trains a one-layer GCN with a small MLP head on the GraphLand benchmark, a collection of industrial graphs with tabular node features.
A simple GCN baseline for node property prediction
PyG's example trains a one-layer GCN with a small MLP head on the GraphLand benchmark, a collection of industrial graphs with tabular node features. GraphLand is not available in K3-Node, so this notebook trains the same model on the Amazon Photo co-purchase graph (7,650 products in 8 categories), with a random 60/20/20 split of the nodes.
Same model as PyG's examples/graphland.py; on Amazon Photo instead of the GraphLand datasets.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
import keras
from keras import ops
from k3_node.datasets import Amazon
from k3_node.layers import GCNConv
from k3_node.loader import FullGraphDataset
from k3_node.transforms import RandomNodeSplit
dataset = Amazon("data/Amazon", name="Photo", transform=RandomNodeSplit(split="train_rest", num_val=0.2, num_test=0.2))
data = dataset[0]
print(data)
Define the model
class Model(keras.Model):
def __init__(self, in_channels, hidden_channels, out_channels):
super().__init__()
self.conv = GCNConv(in_channels, hidden_channels)
self.head = keras.layers.Dense(out_channels)
def call(self, data):
return self.head(ops.relu(self.conv(data.x, data.edge_index)))
model = Model(dataset.num_features, 512, dataset.num_classes)
Train
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.001),
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)