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.
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 nodeGraphClassifier/GraphRegressor: a label or value for every graphLinkPredictor: whether two nodes are linked
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
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"))
Link prediction (Cora)
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.