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