Skip to content

Classifying nodes from their own subgraphs (ShaDow-GNN)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: PubMed (Planetoid)
Description: ShaDow-GNN (Zeng et al., 2021) decouples the depth of a GNN from the size of its receptive field: for every target node it extracts a small subgraph (here the 2-hop neighborhood, with up to 5 neighbors per node) and runs a GNN on that subgraph only.

View in Colab   GitHub source


Classifying nodes from their own subgraphs (ShaDow-GNN)

ShaDow-GNN (Zeng et al., 2021) decouples the depth of a GNN from the size of its receptive field: for every target node it extracts a small subgraph (here the 2-hop neighborhood, with up to 5 neighbors per node) and runs a GNN on that subgraph only. The prediction combines the target node's embedding with the average embedding of its subgraph.

Same model as PyG's examples/shadow.py; on PubMed instead of Flickr, 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. ShaDowKHopSampler yields batches of 1,024 such subgraphs; root_n_id points to the target node of each, and y holds the target nodes' labels.

import keras
from keras import ops
from k3_node.datasets import Planetoid
from k3_node.layers import SAGEConv, global_mean_pool
from k3_node.loader import ShaDowKHopSampler

dataset = Planetoid("data/Planetoid", name="PubMed", split="full")
data = dataset[0]
print(data)

kwargs = dict(depth=2, num_neighbors=5, batch_size=1024)
train_loader = ShaDowKHopSampler(data, node_idx=data.train_mask, shuffle=True, **kwargs)
val_loader = ShaDowKHopSampler(data, node_idx=data.val_mask, **kwargs)
test_loader = ShaDowKHopSampler(data, node_idx=data.test_mask, **kwargs)

Define the model

class GNN(keras.Model):
    def __init__(self, in_channels, hidden_channels, out_channels):
        super().__init__()
        self.convs = [SAGEConv(in_channels, hidden_channels), SAGEConv(hidden_channels, hidden_channels),
                      SAGEConv(hidden_channels, hidden_channels)]
        self.lin = keras.layers.Dense(out_channels)
        self.dropout = keras.layers.Dropout(0.3)

    def call(self, data, training=False):
        x = data.x
        for conv in self.convs:
            x = self.dropout(ops.relu(conv(x, data.edge_index)), training=training)
        # Combine the embedding of each target node with the mean embedding of its subgraph
        root = ops.take(x, data.root_n_id, axis=0)
        return self.lin(ops.concatenate([root, global_mean_pool(x, data.batch, data.num_graphs)], axis=-1))


model = GNN(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),
    metrics=["accuracy"],
)
model.fit(train_loader, validation_data=val_loader, epochs=50, verbose=2)

Evaluate

loss, accuracy = model.evaluate(test_loader, verbose=0)
print(f"Test accuracy: {accuracy:.4f}")