Knowledge graph embeddings: TransE, DistMult, ComplEx and RotatE
Author: K3-Node Team
Backend: Multi-Backend
Dataset: WordNet18RR
Description: Learn an embedding for every entity and relation of a knowledge graph so that true facts score higher than false ones.
Knowledge graph embeddings: TransE, DistMult, ComplEx and RotatE
Learn an embedding for every entity and relation of a knowledge graph so that true facts score higher
than false ones. Each model scores a fact (head, relation, tail) differently, e.g. TransE
(Bordes et al., 2013) treats the relation as a translation:
head + relation ≈ tail. Choose the model with model_name.
Same model as PyG's examples/kge_fb15k_237.py; on WordNet18RR instead of FB15k-237 (a smaller knowledge graph), trained for 200 epochs instead of 500.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
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.
import keras
from k3_node.datasets import WordNet18RR
from k3_node.layers import ComplEx, DistMult, RotatE, TransE
model_name = "transe" # "transe", "complex", "distmult" or "rotate"
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)
Define the model
Every model has its own optimizer settings in PyG's example.
models = {"transe": TransE, "complex": ComplEx, "distmult": DistMult, "rotate": RotatE}
extra_arguments = {"rotate": {"margin": 9.0}}
optimizers = {
"transe": keras.optimizers.Adam(learning_rate=0.01),
"complex": keras.optimizers.Adagrad(learning_rate=0.001, weight_decay=1e-6),
"distmult": keras.optimizers.Adam(learning_rate=0.0001, weight_decay=1e-6),
"rotate": keras.optimizers.Adam(learning_rate=1e-3),
}
model = models[model_name](num_nodes=train_data.num_nodes, num_relations=train_data.num_edge_types,
hidden_channels=50, **extra_arguments.get(model_name, {}))
Train
Each step contrasts 1,000 true facts with corrupted ones (a random head or tail).
model.compile(optimizer=optimizers[model_name])
history = model.fit(train_data, epochs=200, batch_size=1000)
Evaluate
For every test fact, all 40,943 entities are ranked as the possible tail. Reported are the mean rank of the true tail, the mean reciprocal rank (MRR) and how often it is in the top 10.