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