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