Faster neighbor sampling with hierarchical trimming
Author: K3-Node Team
Backend: Multi-Backend
Dataset: PubMed (Planetoid)
Description: With neighbor sampling, the nodes sampled in the last hop only matter for the first GNN layer, the nodes of the second-to-last hop only for the first two layers, and so on.
Faster neighbor sampling with hierarchical trimming
With neighbor sampling, the nodes sampled in the last hop only matter for the first GNN layer, the nodes of the second-to-last hop only for the first two layers, and so on. Hierarchical trimming drops them as soon as they stop mattering, which saves computation without changing the result. This notebook trains a 3-layer GraphSAGE model for one epoch without and with trimming.
Same model as PyG's examples/hierarchical_sampling.py; on PubMed instead of Reddit, to keep the example small.
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. NeighborLoader records how many nodes and edges it sampled in each hop.
import time
import keras
from keras import ops
from k3_node.datasets import Planetoid
from k3_node.loader import NeighborLoader
from k3_node.models import GraphSAGE
dataset = Planetoid("data/Planetoid", name="PubMed", split="full")
data = dataset[0]
print(data)
loader = NeighborLoader(data, input_nodes=data.train_mask, num_neighbors=[20, 10, 5], batch_size=1024, shuffle=True)
Define the model
With trim=True, the model passes the per-hop counts to GraphSAGE, which then trims after every layer.
class Net(keras.Model):
def __init__(self, in_channels, out_channels, trim):
super().__init__()
self.trim = trim
self.sage = GraphSAGE(in_channels, hidden_channels=64, out_channels=out_channels, num_layers=3)
def call(self, data, training=False):
if not self.trim:
return self.sage(data.x, data.edge_index, training=training)
out = self.sage(data.x, data.edge_index, num_sampled_nodes_per_hop=data.num_sampled_nodes,
num_sampled_edges_per_hop=data.num_sampled_edges, training=training)
# The trimmed output has rows for fewer nodes; pad it so every node of the batch has a row
# (only the seed nodes, which come first, count in the loss).
return ops.pad(out, [[0, data.x.shape[0] - out.shape[0]], [0, 0]])
Train with and without trimming
Trimming needs the concrete sizes of every batch, so the models run eagerly.
for trim in [False, True]:
model = Net(dataset.num_features, dataset.num_classes, trim)
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.01),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
weighted_metrics=["accuracy"],
run_eagerly=True,
)
start = time.time()
model.fit(loader, epochs=1, verbose=0)
print(f"trim={trim}: one epoch took {time.time() - start:.1f}s")