Skip to content

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.

View in Colab   GitHub source


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"

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

Evaluate

loss, accuracy = model.evaluate(test_loader, verbose=0)
print(f"Test accuracy: {accuracy:.4f}")