Skip to content

Inductive Deep Graph Infomax with GraphSAGE and neighbor sampling

Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: Learn node embeddings without labels on a graph too large to process at once.

View in Colab   GitHub source


Inductive Deep Graph Infomax with GraphSAGE and neighbor sampling

Learn node embeddings without labels on a graph too large to process at once. The encoder is a 3-layer GraphSAGE network trained on sampled neighborhoods (Hamilton et al., 2017); the objective is Deep Graph Infomax (Veličković et al., 2019), which contrasts real node embeddings with those of a graph whose features were shuffled.

Same model as PyG's examples/infomax_inductive.py; on Cora instead of Reddit, to keep the example small.

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

NeighborLoader samples, for 256 seed nodes at a time, up to 10, 10 and 25 neighbors per hop. The seed nodes come first in every batch.

import keras
from keras import ops
from k3_node.datasets import Planetoid
from k3_node.layers import SAGEConv
from k3_node.loader import NeighborLoader
from k3_node.models import DeepGraphInfomax

dataset = Planetoid("data/Planetoid", name="Cora")
data = dataset[0]
train_loader = NeighborLoader(data, num_neighbors=[10, 10, 25], batch_size=256, shuffle=True)
test_loader = NeighborLoader(data, num_neighbors=[10, 10, 25], batch_size=256)

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. Only the embeddings of the seed nodes of each batch are used.

class Encoder(keras.Model):
    def __init__(self, in_channels, hidden_channels):
        super().__init__()
        self.convs = [SAGEConv(in_channels, hidden_channels), SAGEConv(hidden_channels, hidden_channels),
                      SAGEConv(hidden_channels, hidden_channels)]
        self.activations = [keras.layers.PReLU() for _ in range(3)]

    def call(self, x, edge_index, batch_size=None):
        for conv, act in zip(self.convs, self.activations):
            x = act(conv(x, edge_index))
        return x[:batch_size]  # the seed nodes


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, batch_size=None):
        return keras.random.shuffle(x, seed=self.seed), edge_index, batch_size

    def call(self, data):
        batch_size = ops.shape(data.input_id)[0]  # number of seed nodes
        pos_z, neg_z, summary = self.dgi(data.x, data.edge_index, batch_size=batch_size)
        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.0001))
model.fit(train_loader, epochs=30, 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(test_loader, verbose=0)  # one embedding per node, in node order
train_val_mask = ops.convert_to_numpy(ops.logical_or(data.train_mask, data.val_mask))
test_mask = ops.convert_to_numpy(data.test_mask)
y = ops.convert_to_numpy(data.y)
accuracy = model.dgi.test(z[train_val_mask], y[train_val_mask], z[test_mask], y[test_mask], max_iter=10000)
print(f"Test accuracy: {accuracy:.4f}")