Skip to content

Directed GNN (Dir-GNN) on WebKB

Author: K3-Node Team
Backend: Multi-Backend
Dataset: Texas (WebKB)
Description: Classify web pages of a university website (WebKB) into 5 categories.

View in Colab   GitHub source


Directed GNN (Dir-GNN) on WebKB

Classify web pages of a university website (WebKB) into 5 categories. Links between pages are directed, and on this kind of graph neighbors often belong to different classes. Dir-GNN (Rossi et al., 2023) wraps a GCN so that it aggregates incoming and outgoing links separately and combines them.

Same model as PyG's examples/dir_gnn.py; on WebKB Texas instead of Chameleon to keep the dataset small, trained for 200 epochs instead of 1000.

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

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 WebKB
from k3_node.loader import FullGraphDataset
from k3_node.transforms import NormalizeFeatures

dataset = WebKB("data/WebKB", name="Texas", transform=NormalizeFeatures())
data = dataset[0]
# The dataset comes with 10 different splits; use the first one
data.train_mask, data.val_mask, data.test_mask = data.train_mask[:, 0], data.val_mask[:, 0], data.test_mask[:, 0]
print(data)

Define the model

from k3_node.layers import DirGNNConv, GCNConv


class DirGNN(keras.Model):
    def __init__(self, in_channels, hidden_channels, out_channels, alpha):
        super().__init__()
        self.conv1 = DirGNNConv(GCNConv(in_channels, hidden_channels), alpha, root_weight=False)
        self.conv2 = DirGNNConv(GCNConv(hidden_channels, out_channels), alpha, root_weight=False)

    def call(self, data):
        x = ops.relu(self.conv1(data.x, data.edge_index))
        return self.conv2(x, data.edge_index)


model = DirGNN(dataset.num_features, hidden_channels=128, out_channels=dataset.num_classes, alpha=1.0)

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),
    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}")