Graph classification with a root-node readout (UPFD model)
Author: K3-Node Team
Backend: Multi-Backend
Dataset: PROTEINS (TUDataset)
Description: PyG's UPFD example detects fake news: each graph is a news propagation tree whose root node is the news article itself.
Graph classification with a root-node readout (UPFD model)
PyG's UPFD example detects fake news: each graph is a news propagation tree whose root node is the news article itself. The model runs one GCN layer, max-pools the node embeddings of each graph, and concatenates them with an embedding of the root node. UPFD is not available in K3-Node, so this notebook applies the same model to the PROTEINS dataset (enzyme or not), using the first node of each graph as its "root".
Same model as PyG's examples/upfd.py; on PROTEINS instead of UPFD, which is not available in K3-Node.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
import keras
from keras import ops
from k3_node.datasets import TUDataset
from k3_node.layers import GCNConv, global_max_pool
from k3_node.loader import DataLoader
dataset = TUDataset("data/TU", name="PROTEINS").shuffle()
n = (len(dataset) + 9) // 10 # 10% test, 10% validation, 80% training
test_loader = DataLoader(dataset[:n], batch_size=128)
val_loader = DataLoader(dataset[n:2 * n], batch_size=128)
train_loader = DataLoader(dataset[2 * n:], batch_size=128, shuffle=True)
print(dataset)
Define the model
data.ptr holds where each graph starts, so data.ptr[:-1] are the root nodes.
class RootReadoutNet(keras.Model):
def __init__(self, in_channels, hidden_channels, out_channels):
super().__init__()
self.conv1 = GCNConv(in_channels, hidden_channels)
self.lin0 = keras.layers.Dense(hidden_channels, activation="relu")
self.lin1 = keras.layers.Dense(hidden_channels, activation="relu")
self.lin2 = keras.layers.Dense(out_channels)
def call(self, data):
h = ops.relu(self.conv1(data.x, data.edge_index))
h = global_max_pool(h, data.batch, data.num_graphs)
root = self.lin0(ops.take(data.x, data.ptr[:-1], axis=0)) # the first node of each graph
h = self.lin1(ops.concatenate([root, h], axis=-1))
return self.lin2(h)
model = RootReadoutNet(dataset.num_features, 128, dataset.num_classes)
Train
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.001, weight_decay=0.01),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
metrics=["accuracy"],
)
model.fit(train_loader, validation_data=val_loader, epochs=60, verbose=2)