Skip to content

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.

View in Colab   GitHub source


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"

!pip install k3-node[examples]
import os
os.environ["KERAS_BACKEND"] = "tensorflow"

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)

Evaluate

loss, accuracy = model.evaluate(FullGraphDataset(data, mask="test_mask"), verbose=0)
print(f"Test accuracy: {accuracy:.4f}")