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