Mini-batch training of GNNs and graph transformers
Author: K3-Node Team
Backend: Multi-Backend
Dataset: PubMed (Planetoid)
Description: Train a node classifier on mini-batches sampled with NeighborLoader.
Mini-batch training of GNNs and graph transformers
Train a node classifier on mini-batches sampled with NeighborLoader. Choose the model with
model_name: GraphSAGE, GAT, or the graph transformers SGFormer (Wu et al., 2023)
and Polynormer (Deng et al., 2024). The transformers attend over all nodes of
each seed's sampled neighborhood, so every seed gets its own separate (disjoint) subgraph.
Same model as PyG's examples/ogbn_train.py; on PubMed instead of ogbn-arxiv, to keep the example small, with batches of 256 seed nodes for the graph transformers to fit in CPU memory.
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.loader import FullGraphDataset, NeighborLoader
from k3_node.models import GAT, GraphSAGE, Polynormer, SGFormer
model_name = "sgformer" # "sage", "gat", "sgformer" or "polynormer"
transformer = model_name in ["sgformer", "polynormer"]
dataset = Planetoid("data/Planetoid", name="PubMed", split="full")
data = dataset[0]
print(data)
# Disjoint subgraphs copy every sampled node once per seed, so the transformers use smaller batches
loader_args = dict(num_neighbors=[10, 10, 10], batch_size=256 if transformer else 1024, disjoint=transformer)
train_loader = NeighborLoader(data, input_nodes=data.train_mask, shuffle=True, **loader_args)
val_loader = NeighborLoader(data, input_nodes=data.val_mask, **loader_args)
test_loader = NeighborLoader(data, input_nodes=data.test_mask, **loader_args)
Define the model
def get_model(name, hidden_channels=256, num_layers=3, dropout=0.5):
if name == "gat":
return GAT(dataset.num_features, hidden_channels, num_layers, dataset.num_classes, dropout=dropout, heads=1)
if name == "sage":
return GraphSAGE(dataset.num_features, hidden_channels, num_layers, dataset.num_classes, dropout=dropout)
if name == "sgformer":
return SGFormer(dataset.num_features, hidden_channels, dataset.num_classes, trans_num_heads=1,
trans_dropout=dropout, gnn_num_layers=num_layers, gnn_dropout=dropout)
return Polynormer(dataset.num_features, hidden_channels, dataset.num_classes, local_layers=num_layers)
class Net(keras.Model):
def __init__(self, name):
super().__init__()
self.gnn = get_model(name)
def call(self, data, training=False):
if transformer: # attention within each seed's own subgraph
return self.gnn(data.x, data.edge_index, data.batch, training=training)
return self.gnn(data.x, data.edge_index, training=training)
model = Net(model_name)
Train
Only the seed nodes of each batch count. The graph transformers batch subgraphs of different sizes, so they run eagerly. The learning rate is reduced when the validation accuracy stops improving.
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.003),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
weighted_metrics=["accuracy"],
run_eagerly=transformer,
)
reduce_lr = keras.callbacks.ReduceLROnPlateau(monitor="val_accuracy", mode="max", patience=5)
model.fit(train_loader, validation_data=val_loader, epochs=50, callbacks=[reduce_lr], verbose=2)