Skip to content

Link prediction with graph autoencoders (GAE and VGAE)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: Predict missing edges of the Cora citation graph.

View in Colab   GitHub source


Link prediction with graph autoencoders (GAE and VGAE)

Predict missing edges of the Cora citation graph. A graph autoencoder (Kipf & Welling, 2016) embeds every node with a GCN and scores a pair of nodes by the dot product of their embeddings; it is trained to give high scores to the edges of the graph. Set variational and linear to try the four variants of PyG's example.

Same model as PyG's examples/autoencoder.py; trained for 200 epochs instead of PyG's 400 to keep the notebook quick.

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. Each split has the edges used for message passing (edge_index), the true edges to predict (pos_edge_label_index) and, for validation and testing, as many non-edges (neg_edge_label_index).

import keras
from keras import ops
from k3_node.datasets import Planetoid
from k3_node.layers import GCNConv
from k3_node.models import GAE, VGAE
from k3_node.transforms import Compose, NormalizeFeatures, RandomLinkSplit

variational = False  # VGAE instead of GAE
linear = False       # a single linear GCN layer as encoder

transform = Compose([
    NormalizeFeatures(),
    RandomLinkSplit(num_val=0.05, num_test=0.1, is_undirected=True, split_labels=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

A variational encoder outputs a mean and a log standard deviation per node instead of a single embedding.

class Encoder(keras.Model):
    def __init__(self, in_channels, out_channels, variational, linear):
        super().__init__()
        self.variational, self.linear = variational, linear
        hidden = in_channels if linear else 2 * out_channels
        self.conv1 = None if linear else GCNConv(in_channels, hidden)
        self.conv_mu = GCNConv(hidden, out_channels)
        self.conv_logstd = GCNConv(hidden, out_channels) if variational else None

    def call(self, x, edge_index):
        if not self.linear:
            x = ops.relu(self.conv1(x, edge_index))
        if self.variational:
            return self.conv_mu(x, edge_index), self.conv_logstd(x, edge_index)
        return self.conv_mu(x, edge_index)


encoder = Encoder(dataset.num_features, 16, variational, linear)
model = VGAE(encoder) if variational else GAE(encoder)

Train

model.compile(optimizer=keras.optimizers.Adam(learning_rate=0.01))
history = model.fit(train_data, epochs=200, validation_data=val_data)

Evaluate

scores = model.evaluate(test_data)
print(f"Test AUC: {scores['auc']:.4f}, AP: {scores['ap']:.4f}")