Skip to content

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.

View in Colab   GitHub source


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"

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