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