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