Skip to content

Knowledge graph completion with R-GCN and DistMult

Author: K3-Node Team
Backend: Multi-Backend
Dataset: WordNet18RR
Description: Predict missing facts of a knowledge graph.

View in Colab   GitHub source


Knowledge graph completion with R-GCN and DistMult

Predict missing facts of a knowledge graph. An R-GCN encoder (Schlichtkrull et al., 2018) computes an embedding for every entity from the known facts; a DistMult decoder scores a fact (head, relation, tail) as sum(head * relation * tail). The model learns to score true facts high and corrupted ones low.

Same model as PyG's examples/rgcn_link_pred.py; on WordNet18RR instead of FB15k-237, trained for 200 steps instead of 10,000 (each takes a few seconds on a CPU), and evaluated with the raw tail MRR instead of the filtered one.

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

WordNet18RR links 40,943 English words and meanings (synsets) with 11 relation types such as hypernym or part of. Each fact (head, relation, tail) is an edge with an edge_type; edge_subgraph keeps the facts of each split. The encoder passes messages along both directions of every fact, with a separate relation type for the reverse direction.

import keras
from keras import ops
from k3_node.datasets import WordNet18RR
from k3_node.layers import RGCNConv
from k3_node.training import gradient_step, no_grad

dataset = WordNet18RR("data/WN18RR")
data = dataset[0]
train_data = data.edge_subgraph(data.train_mask)
val_data = data.edge_subgraph(data.val_mask)
test_data = data.edge_subgraph(data.test_mask)
print(train_data)

num_nodes, num_relations = data.num_nodes, train_data.num_edge_types
head, tail = train_data.edge_index[0], train_data.edge_index[1]
edge_index = ops.stack([ops.concatenate([head, tail]), ops.concatenate([tail, head])])
edge_type = ops.concatenate([train_data.edge_type, train_data.edge_type + num_relations])

Define the model

class RGCNEncoder(keras.Model):
    def __init__(self, num_nodes, hidden_channels, num_relations):
        super().__init__()
        self.node_emb = self.add_weight(shape=(num_nodes, hidden_channels), initializer="glorot_uniform")
        self.conv1 = RGCNConv(hidden_channels, hidden_channels, num_relations, num_blocks=5)
        self.conv2 = RGCNConv(hidden_channels, hidden_channels, num_relations, num_blocks=5)
        self.dropout = keras.layers.Dropout(0.2)

    def call(self, edge_index, edge_type, training=False):
        x = ops.relu(self.conv1(self.node_emb, edge_index, edge_type))
        x = self.dropout(x, training=training)
        return self.conv2(x, edge_index, edge_type)


class DistMultDecoder(keras.Model):
    def __init__(self, num_relations, hidden_channels):
        super().__init__()
        self.rel_emb = self.add_weight(shape=(num_relations, hidden_channels), initializer="glorot_uniform")

    def call(self, z, head, relation, tail):
        return ops.sum(ops.take(z, head, axis=0) * ops.take(self.rel_emb, relation, axis=0)
                       * ops.take(z, tail, axis=0), axis=1)


hidden_channels = 500
encoder = RGCNEncoder(num_nodes, hidden_channels, 2 * num_relations)
decoder = DistMultDecoder(num_relations, hidden_channels)

Train

Each step scores all training facts and as many corrupted ones (a random head or tail), with a small penalty on the embedding sizes. Gradients are clipped to norm 1, as in PyG.

def corrupt(head, tail):  # replace either the head or the tail by a random entity
    replace_head = keras.random.uniform(ops.shape(head)) < 0.5
    random_nodes = keras.random.randint(ops.shape(head), 0, num_nodes, dtype=head.dtype)
    return ops.where(replace_head, random_nodes, head), ops.where(replace_head, tail, random_nodes)


def loss_fn():
    z = encoder(edge_index, edge_type, training=True)
    pos = decoder(z, head, train_data.edge_type, tail)
    neg_head, neg_tail = corrupt(head, tail)
    neg = decoder(z, neg_head, train_data.edge_type, neg_tail)
    logits = ops.concatenate([pos, neg])
    labels = ops.concatenate([ops.ones_like(pos), ops.zeros_like(neg)])
    loss = keras.losses.binary_crossentropy(labels, logits, from_logits=True)
    return loss + 1e-2 * (ops.mean(ops.square(z)) + ops.mean(ops.square(decoder.rel_emb)))


optimizer = keras.optimizers.Adam(learning_rate=0.01, global_clipnorm=1.0)
variables = encoder.trainable_variables + decoder.trainable_variables
epochs = 200
for epoch in range(1, epochs + 1):
    loss = gradient_step(loss_fn, variables, optimizer)
    if epoch % 10 == 0:
        print(f"Epoch {epoch:03d}: loss {loss:.4f}")

Evaluate

For every test fact, all entities are ranked as the possible tail by their score; the mean reciprocal rank (MRR) averages 1 / rank of the true tail.

def tail_mrr(data):
    with no_grad():  # evaluation only: nothing to remember for gradients
        z = encoder(edge_index, edge_type)
        h, r, t = data.edge_index[0], data.edge_type, data.edge_index[1]
        scores = ops.matmul(ops.take(z, h, axis=0) * ops.take(decoder.rel_emb, r, axis=0), ops.transpose(z))
        true_scores = ops.take_along_axis(scores, ops.expand_dims(t, 1), axis=1)
        rank = 1 + ops.sum(ops.cast(scores > true_scores, "float32"), axis=1)
    return float(ops.mean(1.0 / rank))


print(f"Validation MRR: {tail_mrr(val_data):.4f}")
print(f"Test MRR: {tail_mrr(test_data):.4f}")