Entity classification with R-GCN (AIFB)
Author: K3-Node Team
Backend: Multi-Backend
Dataset: AIFB (Entities)
Description: Classify the entities of a knowledge graph.
Entity classification with R-GCN (AIFB)
Classify the entities of a knowledge graph. A relational GCN
(Schlichtkrull et al., 2018) uses a separate weight matrix for each
relation type; with num_bases=30, these matrices are combinations of 30 shared ones. The entities
have no features: each one starts from its own learned embedding (x=None).
Same model as PyG's examples/rgcn.py; without its 2-hop subgraph extraction, which only saves memory on larger graphs.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
AIFB is a knowledge graph about a research institute: 8,285 entities linked by 90 relation types (edge_type). The task is to predict the research group of 176 people. The labeled people are given as node indices (train_idx, test_idx) with their labels (train_y, test_y).
import keras
from keras import ops
from k3_node.datasets import Entities
from k3_node.layers import FastRGCNConv
from k3_node.loader import FullGraphDataset
dataset = Entities("data/Entities", "AIFB")
data = dataset[0]
print(data)
Define the model
class RGCN(keras.Model):
def __init__(self, num_nodes, num_classes, num_relations):
super().__init__()
self.conv1 = FastRGCNConv(num_nodes, 16, num_relations, num_bases=30)
self.conv2 = FastRGCNConv(16, num_classes, num_relations, num_bases=30)
def call(self, data):
x = ops.relu(self.conv1(None, data.edge_index, data.edge_type)) # no node features
return self.conv2(x, data.edge_index, data.edge_type)
model = RGCN(data.num_nodes, dataset.num_classes, dataset.num_relations)
Train
FullGraphDataset with index="train_idx" and target="train_y" places each label at its node; only those nodes count in the loss and the accuracy.
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.01, weight_decay=0.0005),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
weighted_metrics=["accuracy"],
)
model.fit(FullGraphDataset(data, index="train_idx", target="train_y"), epochs=50, verbose=2)