Skip to content

Link prediction with a GCN (Cora)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: Predict missing citations in the Cora graph.

View in Colab   GitHub source


Link prediction with a GCN (Cora)

Predict missing citations in the Cora graph. A two-layer GCN embeds every paper; the score of a possible link is the dot product of the two papers' embeddings.

Same model as PyG's examples/link_pred.py.

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

RandomLinkSplit hides 5% of the edges for validation and 10% for testing. Every split keeps the edges used for message passing in edge_index; the edges to predict are in edge_label_index with labels edge_label (1 for a true edge, 0 for a sampled non-edge).

import keras
from keras import ops
from k3_node.datasets import Planetoid
from k3_node.layers import GCNConv
from k3_node.loader import FullGraphDataset
from k3_node.transforms import Compose, NormalizeFeatures, RandomLinkSplit

transform = Compose([
    NormalizeFeatures(),
    RandomLinkSplit(num_val=0.05, num_test=0.1, is_undirected=True, add_negative_train_samples=False),
])
dataset = Planetoid("data/Planetoid", name="Cora", transform=transform)
train_data, val_data, test_data = dataset[0]
print(train_data)

Define the model

class Net(keras.Model):
    def __init__(self, in_channels, hidden_channels, out_channels):
        super().__init__()
        self.conv1 = GCNConv(in_channels, hidden_channels)
        self.conv2 = GCNConv(hidden_channels, out_channels)

    def encode(self, x, edge_index):
        return self.conv2(ops.relu(self.conv1(x, edge_index)), edge_index)

    def decode(self, z, edge_label_index):
        return ops.sum(ops.take(z, edge_label_index[0], axis=0) * ops.take(z, edge_label_index[1], axis=0), axis=-1)

    def call(self, data):
        return self.decode(self.encode(data.x, data.edge_index), data.edge_label_index)


model = Net(dataset.num_features, 128, 64)

Train

neg_sampling_ratio=1.0 adds as many random non-edges as there are training edges, sampled anew every epoch. The score of an edge is a logit, so the loss is binary cross-entropy and the metric is the area under the ROC curve (AUC).

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=0.01),
    loss=keras.losses.BinaryCrossentropy(from_logits=True),
    metrics=[keras.metrics.AUC(from_logits=True, name="auc")],
)
model.fit(
    FullGraphDataset(train_data, neg_sampling_ratio=1.0),
    validation_data=FullGraphDataset(val_data),
    epochs=100,
    verbose=2,
)

Evaluate

loss, auc = model.evaluate(FullGraphDataset(test_data), verbose=0)
print(f"Test AUC: {auc:.4f}")

# Every pair of papers with a positive score is a predicted link
z = model.encode(test_data.x, test_data.edge_index)
print("Predicted links:", int(ops.sum(ops.matmul(z, ops.transpose(z)) > 0)))