Skip to content

Entity classification with relational graph attention (AIFB)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: AIFB (Entities)
Description: Classify the entities of a knowledge graph with relational graph attention (Busbridge et al., 2019): attention weights over neighbors that also depend on the relation type of each edge.

View in Colab   GitHub source


Entity classification with relational graph attention (AIFB)

Classify the entities of a knowledge graph with relational graph attention (Busbridge et al., 2019): attention weights over neighbors that also depend on the relation type of each edge. As in PyG's example, every entity starts from a random 16-dimensional feature vector.

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

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 RGATConv
from k3_node.loader import FullGraphDataset

dataset = Entities("data/Entities", "AIFB")
data = dataset[0]
data.x = keras.random.normal((data.num_nodes, 16))
print(data)

Define the model

class RGAT(keras.Model):
    def __init__(self, in_channels, hidden_channels, out_channels, num_relations):
        super().__init__()
        self.conv1 = RGATConv(in_channels, hidden_channels, num_relations)
        self.conv2 = RGATConv(hidden_channels, hidden_channels, num_relations)
        self.lin = keras.layers.Dense(out_channels)

    def call(self, data):
        x = ops.relu(self.conv1(data.x, data.edge_index, data.edge_type))
        x = ops.relu(self.conv2(x, data.edge_index, data.edge_type))
        return self.lin(x)


model = RGAT(16, 16, 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}")