Skip to content

Mini-batch training with neighbor sampling (GraphSAGE)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: PubMed (Planetoid)
Description: Train a GraphSAGE model (Hamilton et al., 2017) on mini-batches: for every batch of training nodes, NeighborLoader samples up to 25 of their neighbors and 10 of each neighbor's neighbors.

View in Colab   GitHub source


Mini-batch training with neighbor sampling (GraphSAGE)

Train a GraphSAGE model (Hamilton et al., 2017) on mini-batches: for every batch of training nodes, NeighborLoader samples up to 25 of their neighbors and 10 of each neighbor's neighbors. The model never needs the whole graph at once during training, which lets it scale to very large graphs.

Same model as PyG's examples/reddit.py; on PubMed 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

PubMed (19,717 papers, 3 topics) stands in for the much larger graph of PyG's example. With split="full", every paper outside the validation and test sets is a training paper.

import keras
from keras import ops
from k3_node.datasets import Planetoid
from k3_node.layers import SAGEConv
from k3_node.loader import FullGraphDataset, NeighborLoader

dataset = Planetoid("data/Planetoid", name="PubMed", split="full")
data = dataset[0]
print(data)

train_loader = NeighborLoader(data, input_nodes=data.train_mask, num_neighbors=[25, 10], batch_size=1024, shuffle=True)

Define the model

class SAGE(keras.Model):
    def __init__(self, in_channels, hidden_channels, out_channels):
        super().__init__()
        self.conv1 = SAGEConv(in_channels, hidden_channels)
        self.conv2 = SAGEConv(hidden_channels, out_channels)
        self.dropout = keras.layers.Dropout(0.5)

    def call(self, data, training=False):
        x = self.dropout(ops.relu(self.conv1(data.x, data.edge_index)), training=training)
        return self.conv2(x, data.edge_index)


model = SAGE(dataset.num_features, 256, dataset.num_classes)

Train

Only the sampled seed nodes of each batch count in the loss.

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=0.01),
    loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
    weighted_metrics=["accuracy"],
)
model.fit(train_loader, validation_data=FullGraphDataset(data, mask="val_mask"), epochs=10, verbose=2)

Evaluate

PubMed fits in memory, so the trained model is evaluated on the whole graph at once (PyG's example needs layer-wise inference for its much larger graph).

loss, accuracy = model.evaluate(FullGraphDataset(data, mask="test_mask"), verbose=0)
print(f"Test accuracy: {accuracy:.4f}")