Skip to content

Link prediction as graph classification (SEAL)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: SEAL (Zhang & Chen, 2018) predicts whether two papers of the Cora graph are linked by looking only at the structure around them: every candidate link becomes a small graph (the 2-hop neighborhood of both papers) whose nodes are labeled by their distances to the two papers.

View in Colab   GitHub source


Link prediction as graph classification (SEAL)

SEAL (Zhang & Chen, 2018) predicts whether two papers of the Cora graph are linked by looking only at the structure around them: every candidate link becomes a small graph (the 2-hop neighborhood of both papers) whose nodes are labeled by their distances to the two papers. A DGCNN (Zhang et al., 2018) then classifies these graphs.

Same model as PyG's examples/seal_link_pred.py.

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 links for validation and 10% for testing and adds as many non-links. SEALDataset turns every link and non-link into its own graph with label y.

import math

import keras
from keras import ops
from k3_node.datasets import Planetoid, SEALDataset
from k3_node.layers import GCNConv, SortAggregation
from k3_node.loader import DataLoader
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, split_labels=True)
train_data, val_data, test_data = split(dataset[0])

train_dataset = SEALDataset(train_data, num_hops=2)
val_dataset = SEALDataset(val_data, num_hops=2, num_labels=train_dataset.num_labels)
test_dataset = SEALDataset(test_data, num_hops=2, num_labels=train_dataset.num_labels)
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=32)
test_loader = DataLoader(test_dataset, batch_size=32)
print(train_dataset[0])

Define the model

Three GCN layers (plus one with a single channel) compute node features. SortAggregation keeps the k highest-ranked nodes of every graph in a fixed order, so that 1D convolutions can read them like a sequence. k covers 60% of the training graphs.

class DGCNN(keras.Model):
    def __init__(self, in_channels, hidden_channels, num_layers, k):
        super().__init__()
        sizes = [in_channels] + [hidden_channels] * num_layers
        self.convs = [GCNConv(a, b) for a, b in zip(sizes[:-1], sizes[1:])] + [GCNConv(hidden_channels, 1)]
        total_latent_dim = hidden_channels * num_layers + 1
        self.pool = SortAggregation(k)
        self.conv1 = keras.layers.Conv1D(16, total_latent_dim, strides=total_latent_dim, activation="relu")
        self.maxpool1d = keras.layers.MaxPooling1D(2, 2)
        self.conv2 = keras.layers.Conv1D(32, 5, activation="relu")
        self.mlp = keras.Sequential([
            keras.layers.Flatten(),
            keras.layers.Dense(128, activation="relu"),
            keras.layers.Dropout(0.5),
            keras.layers.Dense(1),
        ])

    def call(self, data, training=False):
        xs = [data.x]
        for conv in self.convs:
            xs.append(ops.tanh(conv(xs[-1], data.edge_index)))
        x = ops.concatenate(xs[1:], axis=-1)
        x = self.pool(x, data.batch, dim_size=data.num_graphs)  # [graphs, k * total_latent_dim]
        x = ops.expand_dims(x, -1)  # a sequence with one channel
        x = self.conv2(self.maxpool1d(self.conv1(x)))
        return self.mlp(x, training=training)


num_nodes = sorted(graph.num_nodes for graph in train_dataset)
k = max(10, num_nodes[int(math.ceil(0.6 * len(num_nodes))) - 1])
model = DGCNN(train_dataset.num_features, hidden_channels=32, num_layers=3, k=k)

Train

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=0.0001),
    loss=keras.losses.BinaryCrossentropy(from_logits=True),
    metrics=[keras.metrics.AUC(from_logits=True, name="auc")],
)
model.fit(train_loader, validation_data=val_loader, epochs=50, verbose=2)

Evaluate

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