Skip to content

Link prediction with LPFormer (Cora)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: LPFormer (Shomer et al., 2024) scores a candidate link with a graph transformer: the link attends over the nodes around it (common neighbors, one-hop neighbors and other nodes with a high personalized PageRank from both endpoints), using their PageRank scores as positional encodings, and combines this with the GCN embeddings of the two endpoints.

View in Colab   GitHub source


Link prediction with LPFormer (Cora)

LPFormer (Shomer et al., 2024) scores a candidate link with a graph transformer: the link attends over the nodes around it (common neighbors, one-hop neighbors and other nodes with a high personalized PageRank from both endpoints), using their PageRank scores as positional encodings, and combines this with the GCN embeddings of the two endpoints.

Same model as PyG's examples/lpformer.py; on Cora instead of ogbl-ppa, with the ROC AUC instead of Hits@100.

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 citations for validation and 10% for testing, each with as many non-citations.

import keras
import numpy as np
from keras import ops
from sklearn.metrics import roc_auc_score
from k3_node.datasets import Planetoid
from k3_node.models import LPFormer
from k3_node.training import gradient_step
from k3_node.transforms import RandomLinkSplit

dataset = Planetoid("data/Planetoid", name="Cora")
split = RandomLinkSplit(num_val=0.05, num_test=0.1, is_undirected=True, add_negative_train_samples=False)
train_data, val_data, test_data = split(dataset[0])
train_pos = ops.convert_to_numpy(train_data.edge_label_index)  # the training links (one direction)
print(train_data)

Define the model

The personalized PageRank of every node is computed once, on the training graph.

model = LPFormer(dataset.num_features, 64, num_gnn_layers=3, gnn_dropout=0.1, transformer_dropout=0.1)
ppr = model.calc_sparse_ppr(train_data.edge_index, train_data.num_nodes)

Train

As in PyG, the links being predicted are removed from the graph while they are scored, so the model cannot simply look them up; random node pairs serve as negatives. Cora is small, so batches hold 512 links (PyG uses 32,768 on a graph of millions of links): removing much larger batches would leave almost no graph around the positive links.

optimizer = keras.optimizers.Adam(learning_rate=0.001)
num_nodes = train_data.num_nodes


def train_epoch(batch_size=512):
    losses = []
    perm = np.random.permutation(train_pos.shape[1])
    for start in range(0, len(perm), batch_size):
        batch = perm[start:start + batch_size]
        edges = train_pos[:, batch]
        keep = np.ones(train_pos.shape[1], dtype=bool)
        keep[batch] = False
        kept = train_pos[:, keep]
        masked_graph = np.concatenate([kept, kept[::-1]], axis=1)  # the graph without the batch's links
        neg_edges = np.random.randint(0, num_nodes, size=edges.shape)

        def loss():
            pos_out = model(edges, train_data.x, masked_graph, ppr_matrix=ppr, training=True)
            neg_out = model(neg_edges, train_data.x, train_data.edge_index, ppr_matrix=ppr, training=True)
            return (-ops.mean(ops.log(ops.sigmoid(pos_out) + 1e-6))
                    - ops.mean(ops.log(1 - ops.sigmoid(neg_out) + 1e-6)))

        losses.append(gradient_step(loss, model.trainable_variables, optimizer))
    return float(np.mean(losses))


def auc(data):
    scores = model(data.edge_label_index, data.x, data.edge_index, ppr_matrix=ppr)
    return roc_auc_score(ops.convert_to_numpy(data.edge_label), ops.convert_to_numpy(scores))


model(train_pos[:, :2], train_data.x, train_data.edge_index, ppr_matrix=ppr)  # create the weights
epochs = 100
for epoch in range(1, epochs + 1):
    loss = train_epoch()
    if epoch % 10 == 0:
        print(f"Epoch {epoch:03d}: loss {loss:.4f}, validation AUC {auc(val_data):.4f}")

Evaluate

print(f"Test AUC: {auc(test_data):.4f}")