Mini-batch training on graph clusters (Cluster-GCN)
Author: K3-Node Team
Backend: Multi-Backend
Dataset: PubMed (Planetoid)
Description: Cluster-GCN (Chiang et al., 2019) partitions the graph into many small clusters of densely connected nodes.
Mini-batch training on graph clusters (Cluster-GCN)
Cluster-GCN (Chiang et al., 2019) partitions the graph into many small clusters of densely connected nodes. Each training step uses a few random clusters together with the edges between them, so memory stays small while most edges are kept.
Same model as PyG's examples/cluster_gcn_reddit.py; on PubMed instead of Reddit, to keep the example small (128 clusters instead of 1,500).
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. ClusterData splits it into 128 clusters; every batch combines 20 of them, and only their training nodes count in the loss.
import keras
from keras import ops
from k3_node.datasets import Planetoid
from k3_node.layers import SAGEConv
from k3_node.loader import ClusterData, ClusterLoader, FullGraphDataset
dataset = Planetoid("data/Planetoid", name="PubMed", split="full")
data = dataset[0]
print(data)
cluster_data = ClusterData(data, num_parts=128)
train_loader = ClusterLoader(cluster_data, batch_size=20, shuffle=True).with_mask("train_mask")
Define the model
class Net(keras.Model):
def __init__(self, in_channels, out_channels):
super().__init__()
self.convs = [SAGEConv(in_channels, 128), SAGEConv(128, out_channels)]
self.dropout = keras.layers.Dropout(0.5)
def call(self, data, training=False):
x = self.dropout(ops.relu(self.convs[0](data.x, data.edge_index)), training=training)
return self.convs[1](x, data.edge_index)
model = Net(dataset.num_features, dataset.num_classes)
Train
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.005),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
weighted_metrics=["accuracy"],
)
model.fit(train_loader, validation_data=FullGraphDataset(data, mask="val_mask"), epochs=30, 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).