Forecasting events in a temporal knowledge graph (RE-Net)
Author: K3-Node Team
Backend: Multi-Backend
Dataset: ICEWS18
Description: Predict future events of the ICEWS18 knowledge graph (political events between countries, people and organizations, one graph per day): given a subject and a relation, which object comes next, and vice versa.
Forecasting events in a temporal knowledge graph (RE-Net)
Predict future events of the ICEWS18 knowledge graph (political events between countries, people and organizations, one graph per day): given a subject and a relation, which object comes next, and vice versa. RE-Net (Jin et al., 2020) summarizes the entities each subject (or object) interacted with through the same relation in each of the last 10 days, and reads this history with a GRU.
Same model as PyG's examples/renet.py; on ICEWS18 (PyG's example defaults to the much larger GDELT), trained for 5 epochs instead of 20 (an epoch takes about a quarter of an hour on a CPU).
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
RENet.pre_transform stores, with every event, the recent history of its subject and object (h_sub, h_obj) and the day of each history entry (h_sub_t, h_obj_t). Batching with follow_batch records which event each history entry belongs to.
import keras
import numpy as np
from keras import ops
from k3_node.datasets import ICEWS18
from k3_node.loader import DataLoader
from k3_node.models import RENet
seq_len = 10
pre_transform = RENet.pre_transform(seq_len)
train_dataset = ICEWS18("data/ICEWS18", pre_transform=pre_transform)
test_dataset = ICEWS18("data/ICEWS18", split="test", pre_transform=pre_transform)
train_loader = DataLoader(train_dataset, batch_size=1024, shuffle=True, follow_batch=["h_sub", "h_obj"])
test_loader = DataLoader(test_dataset, batch_size=1024, follow_batch=["h_sub", "h_obj"])
print(train_dataset[0])
Define the model
The model predicts log-probabilities over all entities for the object and for the subject of each event; the wrapper adds both negative log-likelihoods with add_loss.
class EventForecaster(keras.Model):
def __init__(self):
super().__init__()
self.renet = RENet(train_dataset.num_nodes, train_dataset.num_rels, hidden_channels=200, seq_len=seq_len,
dropout=0.5)
def call(self, data, training=False):
log_prob_obj, log_prob_sub = self.renet(data.sub, data.rel, data.obj, data.h_sub, data.h_sub_t,
data.h_sub_batch, data.h_obj, data.h_obj_t, data.h_obj_batch,
training=training)
nll = keras.losses.sparse_categorical_crossentropy
self.add_loss(ops.mean(nll(data.obj, log_prob_obj, from_logits=True))
+ ops.mean(nll(data.sub, log_prob_sub, from_logits=True)))
return log_prob_obj, log_prob_sub
model = EventForecaster()
Train
model.compile(optimizer=keras.optimizers.Adam(learning_rate=0.001, weight_decay=0.00001))
model.fit(train_loader, epochs=5, verbose=2)
Evaluate
Every true object (and subject) is ranked among all entities: mean reciprocal rank (MRR) and how often it is in the top 1, 3 and 10.
results = np.zeros(4)
for i in range(len(test_loader)):
(batch,) = test_loader[i] # a batch in the format the model receives during training
log_prob_obj, log_prob_sub = model(batch)
size = batch.sub.shape[0]
results += ops.convert_to_numpy(model.renet.test(log_prob_obj, batch.obj)) * size
results += ops.convert_to_numpy(model.renet.test(log_prob_sub, batch.sub)) * size
mrr, hits1, hits3, hits10 = results / (2 * len(test_dataset))
print(f"Test MRR: {mrr:.4f}, Hits@1: {hits1:.4f}, Hits@3: {hits3:.4f}, Hits@10: {hits10:.4f}")