Skip to content

LINKX on WebKB

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

View in Colab   GitHub source


LINKX on WebKB

Classify web pages of a university website (WebKB) into 5 categories. On this graph linked pages usually belong to different classes, which hurts standard GNNs. LINKX (Lim et al., 2021) embeds the node features and the adjacency separately with MLPs and then combines them.

Same model as PyG's examples/linkx.py; on WebKB Wisconsin instead of Penn94 to keep the dataset small.

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

dataset = WebKB("data/WebKB", name="Wisconsin")
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.models import LINKX


class LINKXClassifier(keras.Model):
    def __init__(self, num_nodes, in_channels, out_channels):
        super().__init__()
        self.linkx = LINKX(num_nodes, in_channels, hidden_channels=32, out_channels=out_channels,
                           num_layers=1, num_edge_layers=1, num_node_layers=1, dropout=0.5)

    def call(self, data, training=False):
        return self.linkx(data.x, data.edge_index, training=training)


model = LINKXClassifier(data.num_nodes, dataset.num_features, 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=1e-3),
    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}")