Skip to content

Graph U-Net on Cora

Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: Classify papers in the Cora citation network by topic with a Graph U-Net (Gao & Ji, 2019).

View in Colab   GitHub source


Graph U-Net on Cora

Classify papers in the Cora citation network by topic with a Graph U-Net (Gao & Ji, 2019). Like the U-Net used for images, it repeatedly pools the graph down to fewer, more important nodes and then unpools back to the full graph, with skip connections in between. During training, heavy input dropout and random edge dropout (EdgeDropout) regularize the model.

Same model as PyG's examples/graph_unet.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

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

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

Define the model

from k3_node.layers import EdgeDropout
from k3_node.models import GraphUNet


class UNet(keras.Model):
    def __init__(self, in_channels, out_channels, num_nodes):
        super().__init__()
        self.edge_dropout = EdgeDropout(0.2, force_undirected=True)
        self.dropout = keras.layers.Dropout(0.92)
        self.unet = GraphUNet(in_channels, 32, out_channels, depth=3, pool_ratios=[2000 / num_nodes, 0.5])

    def call(self, data, training=False):
        edge_weight = self.edge_dropout(data.edge_index, training=training)  # 0 = dropped edge
        x = self.dropout(data.x, training=training)
        return self.unet(x, data.edge_index, edge_weight=edge_weight)


model = UNet(dataset.num_features, dataset.num_classes, data.num_nodes)

Train

FullGraphDataset feeds the whole graph to Keras; mask selects which nodes count in the loss and the accuracy.

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=0.01, weight_decay=0.001),
    loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
    weighted_metrics=["accuracy"],
)
model.fit(
    FullGraphDataset(data, mask="train_mask"),
    validation_data=FullGraphDataset(data, mask="val_mask"),
    epochs=200,
    verbose=2,
)

Evaluate

loss, accuracy = model.evaluate(FullGraphDataset(data, mask="test_mask"), verbose=0)
print(f"Test accuracy: {accuracy:.4f}")