Scalable Inception Graph Networks (SIGN) on Cora
Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: Classify papers in the Cora citation network by topic.
Scalable Inception Graph Networks (SIGN) on Cora
Classify papers in the Cora citation network by topic. SIGN
(Rossi et al., 2020) does all message passing once, before
training: the SIGN(K=2) transform stores the features averaged over 1-hop and 2-hop neighborhoods
as x1 and x2. The model is then just a set of small networks, one per hop, which makes it fast on
very large graphs.
Same model as PyG's examples/sign.py; on Cora (full graph) instead of Flickr (mini-batches) to keep the dataset small.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
data holds one graph: node features x, edges edge_index, labels y, and masks marking the training, validation and test nodes.
import keras
from keras import ops
from k3_node.datasets import Planetoid
from k3_node.loader import FullGraphDataset
from k3_node.transforms import SIGN, Compose, NormalizeFeatures
dataset = Planetoid("data/Planetoid", name="Cora", transform=Compose([NormalizeFeatures(), SIGN(K=2)]))
data = dataset[0]
print(data)
Define the model
class SIGN(keras.Model):
def __init__(self, out_channels, K):
super().__init__()
self.K = K
self.lins = [keras.layers.Dense(1024, activation="relu") for _ in range(K + 1)]
self.dropout = keras.layers.Dropout(0.5)
self.lin = keras.layers.Dense(out_channels)
def call(self, data, training=False):
xs = [data.x] + [getattr(data, f"x{i}") for i in range(1, self.K + 1)] # 0, 1 and 2 hops
hs = [self.dropout(lin(x), training=training) for x, lin in zip(xs, self.lins)]
return self.lin(ops.concatenate(hs, axis=-1))
model = SIGN(dataset.num_classes, K=2)
Train
FullGraphDataset feeds the whole graph to Keras; mask selects which nodes count in the loss and the accuracy.
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.01),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
weighted_metrics=["accuracy"],
)
model.fit(
FullGraphDataset(data, mask="train_mask"),
validation_data=FullGraphDataset(data, mask="val_mask"),
epochs=200,
verbose=2,
)