Skip to content

Entity classification with R-GCN (AIFB)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: AIFB (Entities)
Description: Classify the entities of a knowledge graph.

View in Colab   GitHub source


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"

!pip install k3-node[examples]
import os
os.environ["KERAS_BACKEND"] = "tensorflow"

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)

Evaluate

loss, accuracy = model.evaluate(FullGraphDataset(data, index="test_idx", target="test_y"), verbose=0)
print(f"Test accuracy: {accuracy:.4f}")