Predicting future interactions with Temporal Graph Networks (TGN)
Author: K3-Node Team
Backend: Multi-Backend
Dataset: mooc (JODIEDataset)
Description: Predict which course activity a student will interact with next.
Predicting future interactions with Temporal Graph Networks (TGN)
Predict which course activity a student will interact with next. A Temporal Graph Network (Rossi et al., 2020) keeps a memory vector for every node that is updated (by a GRU) after each interaction. To score a possible interaction, a graph attention layer combines the memories of the two nodes and of their most recent neighbors.
Same model as PyG's examples/tgn.py; on the MOOC dataset instead of Wikipedia (a 40 MB instead of a 560 MB download).
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
MOOC is a stream of 411,749 timestamped interactions between students and course activities. The events are split in time: the first 70% for training, then 15% each for validation and testing. Each batch of 200 events also gets 200 random destinations as negative examples.
import keras
import numpy as np
from keras import ops
from sklearn.metrics import average_precision_score, roc_auc_score
from k3_node.datasets import JODIEDataset
from k3_node.layers import TransformerConv
from k3_node.loader import TemporalDataLoader
from k3_node.models import IdentityMessage, LastAggregator, LastNeighborLoader, TGNMemory
from k3_node.training import gradient_step, no_grad
data = JODIEDataset("data/JODIE", name="mooc")[0]
train_data, val_data, test_data = data.train_val_test_split(val_ratio=0.15, test_ratio=0.15)
train_loader = TemporalDataLoader(train_data, batch_size=200, neg_sampling_ratio=1.0)
val_loader = TemporalDataLoader(val_data, batch_size=200, neg_sampling_ratio=1.0)
test_loader = TemporalDataLoader(test_data, batch_size=200, neg_sampling_ratio=1.0)
neighbor_loader = LastNeighborLoader(data.num_nodes, size=10) # the 10 latest neighbors of every node
t_all, msg_all = ops.convert_to_numpy(data.t), ops.convert_to_numpy(data.msg)
print(data)
Define the model
class GraphAttentionEmbedding(keras.layers.Layer):
def __init__(self, in_channels, out_channels, msg_dim, time_enc):
super().__init__()
self.time_enc = time_enc
self.conv = TransformerConv(in_channels, out_channels // 2, heads=2, dropout=0.1,
edge_dim=msg_dim + time_enc.out_channels)
def call(self, x, last_update, edge_index, t, msg, training=False):
rel_t = ops.subtract(ops.take(last_update, edge_index[0], axis=0), t) # time since each interaction
edge_attr = ops.concatenate([self.time_enc(ops.cast(rel_t, x.dtype)), msg], axis=-1)
return self.conv(x, edge_index, edge_attr, training=training)
class LinkPredictor(keras.layers.Layer):
def __init__(self, channels):
super().__init__()
self.lin_src = keras.layers.Dense(channels)
self.lin_dst = keras.layers.Dense(channels)
self.lin_final = keras.layers.Dense(1)
def call(self, z_src, z_dst):
return self.lin_final(ops.relu(self.lin_src(z_src) + self.lin_dst(z_dst)))
memory_dim = time_dim = embedding_dim = 100
msg_dim = data.msg.shape[-1]
memory = TGNMemory(data.num_nodes, msg_dim, memory_dim, time_dim,
message_module=IdentityMessage(msg_dim, memory_dim, time_dim),
aggregator_module=LastAggregator())
gnn = GraphAttentionEmbedding(memory_dim, embedding_dim, msg_dim, memory.time_enc)
link_pred = LinkPredictor(embedding_dim)
Train and evaluate
The events are processed in time order. For every batch the model embeds the nodes involved, scores the true and the negative interactions, and then writes the true interactions into the memory and the neighbor lists. assoc maps node ids to rows of the batch's embeddings.
def embed(batch, training):
n_id, edge_index, e_id = neighbor_loader(batch.n_id)
assoc = np.full(data.num_nodes, -1)
assoc[ops.convert_to_numpy(n_id)] = np.arange(n_id.shape[0])
z, last_update = memory(n_id)
e_id = ops.convert_to_numpy(e_id)
z = gnn(z, last_update, edge_index, t_all[e_id], msg_all[e_id], training=training)
src, dst, neg = (assoc[ops.convert_to_numpy(v)] for v in (batch.src, batch.dst, batch.neg_dst))
return (link_pred(ops.take(z, src, axis=0), ops.take(z, dst, axis=0)),
link_pred(ops.take(z, src, axis=0), ops.take(z, neg, axis=0)))
def train_epoch(optimizer, variables):
memory.train()
memory.reset_state() # start with a fresh memory
neighbor_loader.reset_state() # and an empty graph
losses = []
for batch in train_loader:
def loss():
pos_out, neg_out = embed(batch, training=True)
bce = keras.losses.binary_crossentropy
return bce(ops.ones_like(pos_out), pos_out, from_logits=True) + bce(ops.zeros_like(neg_out), neg_out,
from_logits=True)
losses.append(gradient_step(loss, variables, optimizer))
memory.update_state(batch.src, batch.dst, batch.t, batch.msg) # the true interactions
neighbor_loader.insert(batch.src, batch.dst)
return float(np.mean(losses))
def evaluate(loader):
memory.eval()
aps, aucs = [], []
for batch in loader:
with no_grad(): # evaluation only: nothing to remember for gradients
pos_out, neg_out = embed(batch, training=False)
scores = np.concatenate([ops.convert_to_numpy(pos_out), ops.convert_to_numpy(neg_out)])[:, 0]
labels = np.concatenate([np.ones(pos_out.shape[0]), np.zeros(neg_out.shape[0])])
aps.append(average_precision_score(labels, scores))
aucs.append(roc_auc_score(labels, scores))
memory.update_state(batch.src, batch.dst, batch.t, batch.msg)
neighbor_loader.insert(batch.src, batch.dst)
return float(np.mean(aps)), float(np.mean(aucs))
embed(next(iter(train_loader)), training=False) # create the weights
variables = memory.trainable_variables + gnn.trainable_variables + link_pred.trainable_variables
optimizer = keras.optimizers.Adam(learning_rate=0.0001)
epochs = 50
for epoch in range(1, epochs + 1):
loss = train_epoch(optimizer, variables)
val_ap, val_auc = evaluate(val_loader)
test_ap, test_auc = evaluate(test_loader)
print(f"Epoch {epoch:02d}: loss {loss:.4f}, validation AP {val_ap:.4f}, test AP {test_ap:.4f}, test AUC {test_auc:.4f}")