Memory-efficient deep GNNs with reversible layers (RevGNN)
Author: K3-Node Team
Backend: Multi-Backend
Dataset: PubMed (Planetoid)
Description: RevGNN (Li et al., 2021) builds deep GNNs from grouped reversible blocks: the features are split into groups, and each group is updated from the others, so every block's input can be recomputed from its output.
Memory-efficient deep GNNs with reversible layers (RevGNN)
RevGNN (Li et al., 2021) builds deep GNNs from grouped reversible
blocks: the features are split into groups, and each group is updated from the others, so every
block's input can be recomputed from its output. The same dropout mask is shared by all blocks.
Training uses RandomNodeLoader, which splits the nodes into random parts.
Same model as PyG's examples/rev_gnn.py; on PubMed instead of ogbn-products, to keep the example small, trained for 200 epochs instead of PyG's 1,000 to keep the notebook quick.
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. RandomNodeLoader splits it into 10 random parts per epoch; only 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 FullGraphDataset, RandomNodeLoader
from k3_node.models import GroupAddRev
from k3_node.transforms import AddSelfLoops
dataset = Planetoid("data/Planetoid", name="PubMed", split="full", transform=AddSelfLoops())
data = dataset[0]
train_loader = RandomNodeLoader(data, num_parts=10, shuffle=True).with_mask("train_mask")
Define the model
class GNNBlock(keras.layers.Layer):
def __init__(self, in_channels, out_channels):
super().__init__()
self.norm = keras.layers.LayerNormalization()
self.conv = SAGEConv(in_channels, out_channels)
def call(self, x, edge_index, dropout_mask=None):
x = ops.relu(self.norm(x))
if dropout_mask is not None:
x = x * dropout_mask
return self.conv(x, edge_index)
class RevGNN(keras.Model):
def __init__(self, in_channels, hidden_channels, out_channels, num_layers, dropout, num_groups=2):
super().__init__()
self.dropout = dropout
self.seed = keras.random.SeedGenerator(0)
self.lin1 = keras.layers.Dense(hidden_channels)
self.lin2 = keras.layers.Dense(out_channels)
self.norm = keras.layers.LayerNormalization()
group_channels = hidden_channels // num_groups
self.convs = [GroupAddRev([GNNBlock(group_channels, group_channels) for _ in range(num_groups)])
for _ in range(num_layers)]
self.final_dropout = keras.layers.Dropout(dropout)
def call(self, data, training=False):
x = self.lin1(data.x)
mask = None
if training and self.dropout > 0: # one dropout mask, shared by all blocks
keep = keras.random.uniform(ops.shape(x), seed=self.seed) >= self.dropout
mask = ops.cast(keep, x.dtype) / (1 - self.dropout)
for conv in self.convs:
x = conv(x, data.edge_index, mask) if mask is not None else conv(x, data.edge_index)
x = self.final_dropout(ops.relu(self.norm(x)), training=training)
return self.lin2(x)
model = RevGNN(dataset.num_features, hidden_channels=160, out_channels=dataset.num_classes, num_layers=7,
dropout=0.5, num_groups=2)
Train
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.003),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
weighted_metrics=["accuracy"],
)
model.fit(train_loader, validation_data=FullGraphDataset(data, mask="val_mask"), epochs=200, 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).