Zero-shot node classification with RECT on Cora
Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: Classify papers into topics that have no labeled examples at all (zero-shot learning).
Zero-shot node classification with RECT on Cora
Classify papers into topics that have no labeled examples at all (zero-shot learning). RECT (Wang et al., 2020) is trained only on papers from the seen topics: it learns to reproduce a "semantic" description of each seen topic (the average features of its papers). The resulting node embeddings also separate the unseen topics, which we check with a simple logistic-regression classifier.
Same model as PyG's examples/rect.py.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
Features are compressed to 200 dimensions and the graph is smoothed with graph diffusion (GDC). Then topics 1, 2 and 3 are removed from the training labels.
import numpy as np
import keras
from keras import ops
from sklearn.linear_model import LogisticRegression
from k3_node.datasets import Planetoid
from k3_node.loader import FullGraphDataset
from k3_node.models import RECT_L
from k3_node.transforms import GDC, Compose, NormalizeFeatures, RemoveTrainingClasses, SVDFeatureReduction
dataset = Planetoid("data/Planetoid", name="Cora",
transform=Compose([NormalizeFeatures(), SVDFeatureReduction(200), GDC()]))
data = dataset[0]
zs_data = RemoveTrainingClasses([1, 2, 3])(dataset[0]) # training labels only from the seen topics
print(zs_data)
Define the model
The targets are the semantic labels of the training nodes, so the model returns predictions for those nodes only.
class RECT(keras.Model):
def __init__(self, train_index):
super().__init__()
self.rect = RECT_L(200, 200, normalize=False, dropout=0.0)
self.train_index = train_index
def call(self, data, training=False):
out = self.rect(data.x, data.edge_index, data.edge_attr, training=training)
return ops.take(out, self.train_index, axis=0)
model = RECT(train_index=np.where(ops.convert_to_numpy(zs_data.train_mask))[0])
zs_data.y = model.rect.get_semantic_labels(zs_data.x, zs_data.y, zs_data.train_mask)
Train
model.compile(optimizer=keras.optimizers.Adam(learning_rate=0.001, weight_decay=5e-4),
loss=keras.losses.MeanSquaredError(reduction="sum"))
model.fit(FullGraphDataset(zs_data), epochs=200, verbose=0)
Evaluate
Fit a logistic regression on the learned embeddings (using the original labels) and test it on all topics, including the three unseen ones.
embeddings = ops.convert_to_numpy(model.rect.embed(zs_data.x, zs_data.edge_index, zs_data.edge_attr))
y, train_mask, test_mask = (ops.convert_to_numpy(v) for v in (data.y, data.train_mask, data.test_mask))
classifier = LogisticRegression(max_iter=1000).fit(embeddings[train_mask], y[train_mask])
print(f"Test accuracy: {classifier.score(embeddings[test_mask], y[test_mask]):.4f}")