Skip to content

Link prediction with attract-repel embeddings (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 attract-repel embeddings (Cora)

Predict missing citations in the Cora graph. A GCN embeds every paper; a link predictor turns two embeddings into a score. Set use_ar = True for attract-repel embeddings: the first half of each embedding attracts (a large dot product means a likely link) and the second half repels, which can express relations a plain dot product cannot. Otherwise a small MLP scores the pair.

Same model as PyG's examples/ar_link_pred.py; with RandomLinkSplit instead of the deprecated train_test_split_edges.

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 LinkPredictor(keras.Model):  # MLP on the concatenated embeddings
    def __init__(self, hidden_channels):
        super().__init__()
        self.lin1 = keras.layers.Dense(hidden_channels, activation="relu")
        self.lin2 = keras.layers.Dense(1)

    def call(self, z_i, z_j):
        return ops.squeeze(self.lin2(self.lin1(ops.concatenate([z_i, z_j], axis=1))), axis=-1)


class ARLinkPredictor(keras.Model):  # attract-repel score
    def __init__(self, channels):
        super().__init__()
        self.attract_dim = channels // 2

    def call(self, z_i, z_j):
        a = self.attract_dim
        return ops.sum(z_i[:, :a] * z_j[:, :a], axis=1) - ops.sum(z_i[:, a:] * z_j[:, a:], axis=1)


class Net(keras.Model):
    def __init__(self, in_channels, hidden_channels, out_channels, use_ar):
        super().__init__()
        self.conv1 = GCNConv(in_channels, hidden_channels)
        self.conv2 = GCNConv(hidden_channels, out_channels)
        self.predictor = ARLinkPredictor(out_channels) if use_ar else LinkPredictor(hidden_channels)

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

    def call(self, data):
        z = self.encode(data.x, data.edge_index)
        src, dst = data.edge_label_index[0], data.edge_label_index[1]
        return self.predictor(ops.take(z, src, axis=0), ops.take(z, dst, axis=0))


use_ar = True  # attract-repel predictor, or False for the MLP
model = Net(dataset.num_features, 128, 64, use_ar)

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=200,
    verbose=2,
)

Evaluate

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

if use_ar:  # share of the embedding norm spent on repelling
    z = model.encode(test_data.x, test_data.edge_index)
    attract, repel = ops.sum(z[:, :32] ** 2), ops.sum(z[:, 32:] ** 2)
    print(f"R-fraction: {float(ops.convert_to_numpy(repel / (attract + repel))):.4f}")