Skip to content

Unsupervised node embeddings with Deep Graph Infomax (Cora)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: Learn node embeddings of the Cora citation graph without using any labels.

View in Colab   GitHub source


Unsupervised node embeddings with Deep Graph Infomax (Cora)

Learn node embeddings of the Cora citation graph without using any labels. Deep Graph Infomax (Veličković et al., 2019) maximizes the mutual information between node embeddings and a summary of the whole graph.

Same model as PyG's examples/infomax_transductive.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

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.models import DeepGraphInfomax

dataset = Planetoid("data/Planetoid", name="Cora")
data = dataset[0]

Define the model

Deep Graph Infomax trains the encoder to tell the real graph apart from a corrupted one, whose node features are shuffled among the nodes. The summary of the graph is the sigmoid of the mean node embedding. The wrapper adds DGI's loss with add_loss, so Keras' fit can train it without labels.

class Encoder(keras.Model):
    def __init__(self, in_channels, hidden_channels):
        super().__init__()
        self.conv = GCNConv(in_channels, hidden_channels)
        self.prelu = keras.layers.PReLU()

    def call(self, x, edge_index):
        return self.prelu(self.conv(x, edge_index))


class DGI(keras.Model):
    def __init__(self, in_channels, hidden_channels):
        super().__init__()
        self.seed = keras.random.SeedGenerator(0)
        self.dgi = DeepGraphInfomax(
            hidden_channels,
            encoder=Encoder(in_channels, hidden_channels),
            summary=lambda z, *args, **kwargs: ops.sigmoid(ops.mean(z, axis=0)),
            corruption=self.corruption,
        )

    def corruption(self, x, edge_index):  # shuffle the node features
        return keras.random.shuffle(x, seed=self.seed), edge_index

    def call(self, data):
        pos_z, neg_z, summary = self.dgi(data.x, data.edge_index)
        self.add_loss(self.dgi.loss(pos_z, neg_z, summary))
        return pos_z


model = DGI(dataset.num_features, 512)

Train

model.compile(optimizer=keras.optimizers.Adam(learning_rate=0.001))
model.fit(FullGraphDataset(data), epochs=300, verbose=2)

Evaluate

The embeddings were learned without labels. A logistic regression trained on the embeddings of the training nodes shows how much class information they contain.

z = model.predict(FullGraphDataset(data), verbose=0)
accuracy = model.dgi.test(z[data.train_mask], data.y[data.train_mask], z[data.test_mask], data.y[data.test_mask],
                          max_iter=150)
print(f"Test accuracy: {accuracy:.4f}")