Mini-batch training on sampled subgraphs (GraphSAINT)
Author: K3-Node Team
Backend: Multi-Backend
Dataset: PubMed (Planetoid)
Description: GraphSAINT (Zeng et al., 2020) trains on small subgraphs sampled from the full graph: here, the nodes visited by short random walks.
Mini-batch training on sampled subgraphs (GraphSAINT)
GraphSAINT (Zeng et al., 2020) trains on small subgraphs sampled from the full graph: here, the nodes visited by short random walks. The model is a 3-layer GNN whose layer outputs are concatenated before the classifier (jumping knowledge).
Same model as PyG's examples/graph_saint.py; on PubMed instead of Flickr, to keep the example small (without the optional GraphSAINT loss normalization).
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. Each of the 5 steps per epoch samples random walks of length 2 from 1,500 start nodes.
import keras
from keras import ops
from k3_node.datasets import Planetoid
from k3_node.layers import GraphConv
from k3_node.loader import FullGraphDataset, GraphSAINTRandomWalkSampler
dataset = Planetoid("data/Planetoid", name="PubMed", split="full")
data = dataset[0]
print(data)
loader = GraphSAINTRandomWalkSampler(data, batch_size=1500, walk_length=2, num_steps=5).with_mask("train_mask")
Define the model
class Net(keras.Model):
def __init__(self, in_channels, hidden_channels, out_channels):
super().__init__()
self.conv1 = GraphConv(in_channels, hidden_channels, aggr="mean")
self.conv2 = GraphConv(hidden_channels, hidden_channels, aggr="mean")
self.conv3 = GraphConv(hidden_channels, hidden_channels, aggr="mean")
self.lin = keras.layers.Dense(out_channels)
self.dropout = keras.layers.Dropout(0.2)
def call(self, data, training=False):
x1 = self.dropout(ops.relu(self.conv1(data.x, data.edge_index)), training=training)
x2 = self.dropout(ops.relu(self.conv2(x1, data.edge_index)), training=training)
x3 = self.dropout(ops.relu(self.conv3(x2, data.edge_index)), training=training)
return self.lin(ops.concatenate([x1, x2, x3], axis=-1))
model = Net(dataset.num_features, 256, dataset.num_classes)
Train
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.001),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
weighted_metrics=["accuracy"],
)
model.fit(loader, validation_data=FullGraphDataset(data, mask="val_mask"), epochs=50, 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).