-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathevaluate.py
More file actions
39 lines (30 loc) · 1.35 KB
/
Copy pathevaluate.py
File metadata and controls
39 lines (30 loc) · 1.35 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
import dgl
import numpy as np
import torch
from tqdm import tqdm
from utils.metrics import compute_link_prediction_metrics, get_ranks_from_scores
@torch.no_grad()
def evaluate(model, data_loader, G, static_emb, dyn_entity_emb, dyn_relation_emb,
num_rels, cfg, split="Validation"):
model.eval()
device = cfg.runtime.device
all_ranks = []
for prior_G, batch_G, cumul_G, batch_t in tqdm(data_loader, desc=f"[{split}]"):
batch_G = batch_G.to(device)
nids = batch_G.ndata['_ID'].long().cpu()
static_e = static_emb.structural[nids].to(device)
dynamic_e = dyn_entity_emb.structural[nids, -1].to(device)
combined = model.combiner(static_e, dynamic_e)
dyn_rel = dyn_relation_emb.structural[:, -1]
_, _, _, tail_pred = model.edge_model(
batch_G, combined, static_e, dynamic_e, dyn_rel, return_pred=True,
)
heads, tails = batch_G.edges()
target_global = batch_G.ndata[dgl.NID][tails.long()].long()
ranks = get_ranks_from_scores(tail_pred, target_global)
all_ranks.extend(ranks.tolist())
dyn_entity_emb, dyn_relation_emb = model.embedding_updater(
prior_G, batch_G, cumul_G, static_emb, dyn_entity_emb, dyn_relation_emb, device,
)
metrics = compute_link_prediction_metrics(np.array(all_ranks))
return metrics