Skip to content

High-level task APIs

Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: K3-Node's task estimators train, evaluate and predict on graphs in a few lines, in the style of scikit-learn.

View in Colab   GitHub source


High-level task APIs

K3-Node's task estimators train, evaluate and predict on graphs in a few lines, in the style of scikit-learn. Each one wraps a graph neural network (the backbone, e.g. "gcn", "gat", "sage" or "gin") and chooses a suitable loss, readout and metrics:

  • NodeClassifier / NodeRegressor: a label or value for every node
  • GraphClassifier / GraphRegressor: a label or value for every graph
  • LinkPredictor: whether two nodes are linked

Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"

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

Node classification (Cora)

Predict the topic of every paper in the Cora citation graph.

from k3_node.datasets import MoleculeNet, Planetoid, TUDataset
from k3_node.tasks import GraphClassifier, GraphRegressor, LinkPredictor, NodeClassifier

cora = Planetoid("data/Planetoid", name="Cora")[0]

classifier = NodeClassifier(backbone="gcn", hidden_channels=64, num_layers=2, dropout=0.5)
classifier.fit(cora, epochs=100, lr=0.01, verbose=0)
print(classifier.evaluate(cora, mask="test_mask"))

Predict which pairs of papers cite each other. RandomLinkSplit holds out 10% of the citations (and as many non-citations) for testing.

from k3_node.transforms import RandomLinkSplit

train_data, val_data, test_data = RandomLinkSplit(num_val=0.05, num_test=0.1, is_undirected=True)(cora)

link_predictor = LinkPredictor(backbone="gcn", hidden_channels=64, out_channels=32, decoder="inner_product")
link_predictor.fit(train_data, epochs=100, lr=0.01, verbose=0)
print(link_predictor.evaluate(test_data))

Graph classification (MUTAG)

Predict whether a molecule is mutagenic.

mutag = TUDataset("data/TU", name="MUTAG").shuffle()

graph_classifier = GraphClassifier(backbone="gin", hidden_channels=64, num_layers=3, pooling="mean", dropout=0.5)
graph_classifier.fit(mutag[:150], epochs=50, batch_size=32, lr=0.01, verbose=0)
print(graph_classifier.evaluate(mutag[150:], batch_size=32))

Graph regression (ESOL)

Predict the water solubility of molecules.

esol = MoleculeNet("data/MoleculeNet", "ESOL").shuffle()

regressor = GraphRegressor(backbone="sage", hidden_channels=64, num_layers=3, pooling="mean")
regressor.fit(esol[:1000], epochs=50, batch_size=64, verbose=0)
print(regressor.evaluate(esol[1000:], batch_size=64))