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