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