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