Skip to content

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

View in Colab   GitHub source


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"

!pip install k3-node[examples]
import os
os.environ["KERAS_BACKEND"] = "tensorflow"

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