Skip to content

Distilling a GNN into an MLP (GLNN)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: Graph-less neural networks (Zhang et al., 2021): train a GCN (the teacher), then train a plain MLP (the student) to reproduce the teacher's predictions from the node features alone.

View in Colab   GitHub source


Distilling a GNN into an MLP (GLNN)

Graph-less neural networks (Zhang et al., 2021): train a GCN (the teacher), then train a plain MLP (the student) to reproduce the teacher's predictions from the node features alone. The student no longer needs the graph at prediction time, which makes it much faster.

Same model as PyG's examples/glnn.py.

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

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 NormalizeFeatures
from k3_node.models import GCN, MLP
from k3_node.training import gradient_step

dataset = Planetoid("data/Planetoid", name="Cora", transform=NormalizeFeatures())
data = dataset[0]
print(data)

Train the teacher

class Teacher(keras.Model):
    def __init__(self):
        super().__init__()
        self.gcn = GCN(dataset.num_features, hidden_channels=16, out_channels=dataset.num_classes, num_layers=2)

    def call(self, data, training=False):
        return self.gcn(data.x, data.edge_index, training=training)


teacher = Teacher()
teacher.compile(
    optimizer=keras.optimizers.Adam(learning_rate=0.01, weight_decay=5e-4),
    loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
    weighted_metrics=["accuracy"],
)
teacher.fit(FullGraphDataset(data, mask="train_mask"), epochs=200, verbose=0)
_, accuracy = teacher.evaluate(FullGraphDataset(data, mask="test_mask"), verbose=0)
print(f"Teacher test accuracy: {accuracy:.4f}")

Train the student

The student's loss mixes the usual cross-entropy on the training labels (weight lamb) with the KL divergence to the teacher's predictions on all nodes (weight 1 - lamb).

lamb = 0.0  # PyG's default: learn from the teacher only
soft_labels = ops.log_softmax(teacher.predict(FullGraphDataset(data), verbose=0), axis=-1)
student = MLP([dataset.num_features, 64, dataset.num_classes], dropout=0.5, norm=None)
optimizer = keras.optimizers.Adam(learning_rate=0.01, weight_decay=5e-4)
train_weight = ops.cast(data.train_mask, "float32")


def student_loss():
    log_probs = ops.log_softmax(student(data.x, training=True), axis=-1)
    hard = keras.losses.sparse_categorical_crossentropy(data.y, log_probs, from_logits=True)
    hard = ops.sum(hard * train_weight) / ops.sum(train_weight)
    soft = ops.mean(ops.sum(ops.exp(soft_labels) * (soft_labels - log_probs), axis=-1))  # KL(teacher || student)
    return lamb * hard + (1 - lamb) * soft


student(data.x)  # create the weights
for epoch in range(1, 501):
    loss = gradient_step(student_loss, student.trainable_variables, optimizer)
    if epoch % 100 == 0:
        print(f"Epoch {epoch:03d}: loss {loss:.4f}")

Evaluate

predictions = ops.argmax(student(data.x), axis=-1)
correct = ops.cast(predictions == ops.cast(data.y, predictions.dtype), "float32")
test = ops.cast(data.test_mask, "float32")
print(f"Test accuracy: {float(ops.sum(correct * test) / ops.sum(test)):.4f}")