Coverage for hopwise/model/knowledge_aware_recommender/cafe.py: 83%
663 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
1# @Time : 2025/02/19
2# @Author : Alessandro Soccol
3# @Email : alessandro.soccol@unica.it
5"""CAFE
6##################################################
7Reference:
8 Xian et al. "CAFE: Coarse-to-Fine Neural Symbolic Reasoning for Explainable Recommendation." in CIKM 2020.
10Reference code:
11 https://github.com/orcax/CAFE
12"""
14import random
15from functools import reduce
17import numpy as np
18import torch
19import torch.nn.functional as F
20from torch import nn
22from hopwise.model.abstract_recommender import KnowledgeRecommender
23from hopwise.utils import InputType
26class CAFE(KnowledgeRecommender):
27 """
28 CAFE is a knowledge-aware recommender system that uses symbolic reasoning
29 over a knowledge graph to explain recommendations.
31 Note:
32 Assumes that each relation corresponds to a unique pair of entity types. e.g. ui-relation -> (user, item)
33 """
35 input_type = InputType.USERWISE
37 def __init__(self, config, dataset):
38 super().__init__(config, dataset)
39 self.dataset = dataset
41 # Load parameters info from config
42 self.device = config["device"]
43 self.load_embeddings = config["load_embeddings"]
44 self.raw_metapaths = config["path_constraint"]
46 # Load CAFE parameters
47 self.rank_weight = config["rank_weight"]
48 self.deep_module = config["deep_module"]
49 self.use_dropout = config["use_dropout"]
50 self.topk_candidates = config["topk_candidates"]
51 self.sample_size = config["sample_size"]
52 self.topk_paths = config["topk_paths"]
53 self.path_max_user_trials = config["max_user_trials"]
55 # user-item relation
56 self.ui_relation = dataset.ui_relation
57 self.ui_relation_id = dataset.field2token_id["relation_id"][dataset.ui_relation]
59 # Load Knowledge Graph Embedding Checkpoint
60 self.user_embedding = dataset.get_preload_weight("user_embedding_id")
61 self.entity_embedding = dataset.get_preload_weight("entity_embedding_id")
62 self.relation_embedding = dataset.get_preload_weight("relation_embedding_id")
64 # Topk Candidates
65 self.topk_user_items = self._compute_top_items()
66 # Turn into torch, so that the weight is updated.
67 self.user_embedding = torch.from_numpy(self.user_embedding).to(device=self.device, dtype=torch.float32)
68 self.entity_embedding = torch.from_numpy(self.entity_embedding).to(self.device, dtype=torch.float32)
69 self.relation_embedding = torch.from_numpy(self.relation_embedding).to(self.device, dtype=torch.float32)
70 self.embedding_size = self.user_embedding.size(1)
72 # Embedding mapping
73 self.embeddings = {
74 "user": self.user_embedding,
75 "entity": self.entity_embedding,
76 "relation": self.relation_embedding,
77 }
79 # Map Relation ID to relation name to check has pattern constraint
80 self.rid2relation = {v: k for k, v in dataset.field2token_id["relation_id"].items()}
81 self.relation2rid = dataset.field2token_id["relation_id"]
83 # Positives
84 self.positives = dataset.history_item_matrix()[0]
86 # Load Full Knowledge Graph in dict form
87 self.graph_dict = dataset.ckg_dict_graph(ui_bidirectional=False)
88 self.memory_size = 10000 # number of paths to save for each metapath
89 self.replay_memory = {}
91 self.relation_info = dict()
92 for relation_name in self.rid2relation.values():
93 if relation_name == "[PAD]":
94 continue
96 if relation_name == f"{dataset.ui_relation}_r":
97 raise ValueError("The ui_relation name should not end with '_r'.")
99 if relation_name == dataset.ui_relation:
100 head, tail = "user", "entity"
101 else:
102 head, tail = "entity", "entity"
104 if relation_name not in self.relation_info:
105 self.relation_info[relation_name] = {"name": relation_name, "entity_head": head, "entity_tail": tail}
107 # Transform each node in the metapath in a tuple
108 self.metapaths = []
109 for mp in self.raw_metapaths:
110 metapath = []
111 for node in mp:
112 metapath.append((node[0], node[1]))
113 self.metapaths.append(metapath)
115 self.mpath_ids = list(range(len(self.metapaths)))
117 for mpid in range(len(self.metapaths)):
118 self.replay_memory[mpid] = ReplayMemory(self.memory_size)
120 self.model = SymbolicNetwork(
121 self.relation_info,
122 self.relation2rid,
123 self.embeddings,
124 self.embedding_size,
125 self.deep_module,
126 self.use_dropout,
127 self.n_items,
128 self.device,
129 )
131 # random generator
132 self.rng = np.random.default_rng()
134 def _compute_top_items(self):
135 u_p_scores = np.dot(
136 self.user_embedding + self.relation_embedding[self.ui_relation_id], self.entity_embedding[: self.n_items].T
137 )
138 ui_scores = np.argsort(u_p_scores, axis=1) # From worst to best
139 top100_ui_scores = ui_scores[:, -100:][:, ::-1]
140 topk_user_items = top100_ui_scores[:, : self.topk_candidates]
141 return topk_user_items
143 def _get_batch_by_user(self, users):
144 pos_path_batch, neg_pid_batch = [], []
145 top_pids = np.arange(self.topk_candidates)
146 user_trials = {u.item(): 0 for u in users} # Track trials per user
147 skipped_users = []
148 # it select the it user, if a path is not found, another metapath and item is used.
149 it = 0
150 while len(pos_path_batch) < len(users) and it < len(users):
151 # Take the user
152 user = users[it].item()
154 # Skip if max trials exceeded
155 if user_trials[user] >= self.path_max_user_trials:
156 if user not in skipped_users:
157 skipped_users.append(user)
158 it += 1
159 continue
160 # Sample a metapath
161 mpid = self.rng.choice(self.mpath_ids)
163 # Sample one of the best topk_candidates e.g. sample one of the top20 pids
164 pidx = self.rng.choice(top_pids)
166 # Take the corresponding score
167 item = self.topk_user_items[user][pidx]
169 # Compute the probability to sample path from memory, P \in [0, 0.5].
170 use_memory_prob = 0.5 * len(self.replay_memory[mpid]) / self.memory_size
172 # Sample a history path from memory.
173 if self.rng.random() < use_memory_prob:
174 hist_path = self.replay_memory[mpid].sample()
175 pos_path_batch.append(hist_path)
176 # Sample a new path from graph.
177 else:
178 paths = self.fast_sample_path_with_target(mpid, user, item, 1)
179 # if a path is not found, try again
180 if not paths:
181 user_trials[user] += 1
182 continue
183 pos_path_batch.append(paths[0])
184 self.replay_memory[mpid].add(paths)
186 # Sample a negative item.
187 if pidx < self.topk_candidates - 1:
188 neg_pidx = self.rng.choice(np.arange(pidx + 1, self.topk_candidates))
189 neg_pid = self.topk_user_items[user][neg_pidx]
190 else:
191 neg_pid = self.rng.choice(self.n_items)
193 neg_pid_batch.append(neg_pid)
194 it += 1
196 pos_path_batch = np.array(pos_path_batch)
197 neg_pid_batch = np.array(neg_pid_batch)
198 return mpid, pos_path_batch, neg_pid_batch
200 def _rev_rel(self, rel):
201 if rel == self.ui_relation:
202 return self.ui_relation
204 if rel.endswith("_r"):
205 return rel[:-2]
206 return rel + "_r"
208 def fast_sample_path_with_target(self, mpath_id, user_id, target_id, num_paths, sample_size=100):
209 """Sample one path given source and target, using BFS from both sides.
211 Returns:
212 list: List of entity ids forming the path.
213 """
214 metapath = self.metapaths[mpath_id]
215 path_len = len(metapath) - 1
216 mid_level = int((path_len + 0) / 2)
218 # Forward BFS (e.g. u--e1--e2--e3).
219 forward_paths = [[user_id]]
220 for i in range(1, mid_level + 1):
221 _, last_entity = metapath[i - 1]
222 next_relation, _ = metapath[i]
223 tmp_paths = []
224 for fp in forward_paths:
225 try:
226 next_ids = self.graph_dict[last_entity][fp[-1]][self.relation2rid[next_relation]]
227 except KeyError:
228 next_ids = []
229 # Random sample ids
230 if len(next_ids) > sample_size:
231 # next_ids = np.random.permutation(next_ids)[:sample_size]
232 next_ids = self.rng.choice(next_ids, size=sample_size, replace=False)
233 for next_id in next_ids:
234 tmp_paths.append(fp + [next_id])
235 forward_paths = tmp_paths
236 # Backward BFS (e.g. e4--e5--e6).
237 backward_paths = [[target_id]]
238 for i in reversed(range(mid_level + 2, path_len + 1)): # i=l, l-2,..., mid+2
239 next_relation, next_entity = metapath[i]
240 tmp_paths = []
241 for bp in backward_paths:
242 try:
243 curr_ids = self.graph_dict[next_entity][bp[0]][self.relation2rid[self._rev_rel(next_relation)]]
244 except KeyError:
245 curr_ids = []
246 # Random sample ids
247 if len(curr_ids) > sample_size:
248 curr_ids = self.rng.choice(curr_ids, size=sample_size, replace=False)
249 for curr_id in curr_ids:
250 tmp_paths.append([curr_id] + bp)
251 backward_paths = tmp_paths
252 # Build hash map for indexing backward paths.
253 # e.g. a dict with key=e3 and value=(e4--e5--e6).
254 backward_map = {}
255 next_relation, next_entity = metapath[mid_level + 1]
256 # convert relation name to id for graphdict lookup
257 for bp in backward_paths:
258 try:
259 curr_ids = self.graph_dict[next_entity][bp[0]][self.relation2rid[self._rev_rel(next_relation)]]
260 except KeyError:
261 try:
262 relations = list(self.graph_dict[last_entity][fp[-1]].keys())
263 curr_ids = self.graph_dict[last_entity][fp[-1]][self.rng.choice(relations)]
264 except KeyError:
265 curr_ids = []
266 if len(curr_ids) > sample_size:
267 curr_ids = self.rng.choice(curr_ids, size=sample_size, replace=False)
268 for curr_id in curr_ids:
269 if curr_id not in backward_map:
270 backward_map[curr_id] = []
271 backward_map[curr_id].append(bp)
272 # Find intersection of forward paths and backward paths.
273 final_paths = []
274 for fp_idx in self.rng.permutation(len(forward_paths)):
275 fp = forward_paths[fp_idx]
276 mid_id = fp[-1]
277 if mid_id not in backward_map:
278 continue
279 self.rng.shuffle(backward_map[mid_id])
280 for bp in backward_map[mid_id]:
281 final_paths.append(fp + bp)
282 if len(final_paths) >= num_paths:
283 break
284 if len(final_paths) >= num_paths:
285 break
287 return final_paths
289 def count_paths_with_target(self, mpath_id, user_id, target_id, sample_size=50):
290 """This is an approx count, not exact."""
291 if isinstance(target_id, torch.Tensor):
292 target_id = target_id.item()
293 metapath = self.metapaths[mpath_id]
294 path_len = len(metapath) - 1
295 mid_level = int((path_len + 0) / 2)
297 # Forward BFS (e.g. u--e1--e2--e3).
298 forward_ids = [user_id]
299 for i in range(1, mid_level + 1): # i=1, 2,..., mid
300 _, last_entity = metapath[i - 1]
301 next_relation, _ = metapath[i]
302 tmp_ids = []
303 for eid in forward_ids:
304 try:
305 next_ids = self.graph_dict[last_entity][eid][self.relation2rid[next_relation]]
306 if len(next_ids) > sample_size:
307 next_ids = self.rng.choice(next_ids, size=sample_size, replace=False).tolist()
308 tmp_ids.extend(next_ids)
309 except KeyError:
310 continue
311 forward_ids = tmp_ids
313 # Backward BFS (e.g. e4--e5--e6).
314 backward_ids = [target_id]
315 for i in reversed(range(mid_level + 1, path_len + 1)): # i=l, l-1,..., mid+1
316 next_relation, next_entity = metapath[i]
317 tmp_ids = []
318 for eid in backward_ids:
319 try:
320 curr_ids = self.graph_dict[next_entity][eid][self.relation2rid[self._rev_rel(next_relation)]]
321 except KeyError:
322 curr_ids = []
323 tmp_ids.extend(curr_ids)
324 backward_ids = tmp_ids
326 count = len(set(forward_ids).intersection(backward_ids))
327 return count
329 def calculate_loss(self, interaction):
330 users = interaction[self.USER_ID]
331 users = users[users != 0]
333 mpid, pos_paths, neg_pids = self._get_batch_by_user(users)
335 pos_paths = torch.from_numpy(pos_paths).to(self.device)
336 neg_pids = torch.from_numpy(neg_pids).to(self.device)
337 reg_loss, rank_loss = self.model.forward(self.metapaths[mpid], pos_paths, neg_pids)
338 rank_loss *= self.rank_weight
340 return reg_loss, rank_loss
342 def predict(self, interaction):
343 return
345 def full_sort_predict(self, interaction):
346 users = interaction[self.USER_ID]
347 kg_mask = KGMask(self.graph_dict, self.ui_relation_id)
348 predicted_paths = self._infer_paths(users, kg_mask)
349 path_counts = self._estimate_path_count(users)
350 results = self.run_program(users, path_counts, predicted_paths)
351 scores, paths = results
352 paths = self.convert_path_relations(paths)
353 return scores, paths
355 def convert_path_relations(self, paths):
356 new_data = []
357 for user, item, score, path in paths:
358 sanitized_path = [path[0]] + [
359 (self.relation2rid[relation], e_type, eid) for relation, e_type, eid in path[1:]
360 ]
361 new_data.append([user, item, score, sanitized_path])
362 return new_data
364 def explain(self, interaction):
365 """Support function used for case study.
367 Args:
368 interaction : test interaction data
370 Returns:
371 pd.Dataframe: explanation results with columns: "user", "item", "score", "path"
372 """
373 users = interaction[self.USER_ID]
375 kg_mask = KGMask(self.graph_dict, self.ui_relation_id)
376 predicted_paths = self._infer_paths(users, kg_mask)
377 path_counts = self._estimate_path_count(users)
379 scores, explanations = self.run_program(users, path_counts, predicted_paths)
381 for exp in explanations:
382 exp[-1] = self.decode_path(exp[-1])
384 return scores, explanations
386 def decode_path(self, path):
387 return path
389 def _infer_paths(self, users, kg_mask):
390 predictions = dict()
391 for user in users:
392 user = user.item() # noqa: PLW2901
394 predictions[user] = dict()
395 for mpid, metapath in enumerate(self.metapaths):
396 paths = self.model.infer_with_path(
397 metapath, user, kg_mask, excluded_pids=self.positives[user], topk_paths=self.topk_paths
398 )
399 predictions[user][mpid] = paths
400 return predictions
402 def _estimate_path_count(self, users):
403 num_mp = len(self.metapaths)
404 counts = dict()
405 for user in users:
406 user = user.item() # noqa: PLW2901
408 counts[user] = np.zeros(num_mp)
409 positives = self.positives[user]
410 positives = positives[positives != 0] # due to padding
411 for pid in positives:
412 for mpid in range(num_mp):
413 cnt = self.count_paths_with_target(mpid, user, pid, 50)
414 counts[user][mpid] += cnt
415 counts[user] = counts[user] / len(self.positives[user])
416 return counts
418 def run_program(self, users, path_counts, predicted_paths):
419 results = torch.full((len(users), self.n_items), -torch.inf)
420 collect_results = list()
422 kg_mask = KGMask(self.graph_dict, self.ui_relation_id)
423 program_exe = MetaProgramExecutor(self.model, self.rng, self.device, kg_mask, self.relation2rid)
425 pred_paths_instances = dict()
426 for i, user in enumerate(users):
427 user = user.item() # noqa: PLW2901
429 pred_paths_instances[user] = dict()
430 program = self.create_heuristic_program(self.metapaths, predicted_paths[user], path_counts[user])
431 positives = self.positives[user]
432 positives = positives[positives != 0] # due to padding
433 program_exe.execute(program, user, positives)
434 paths = program_exe.collect_results(program)
435 tmp = [(r[0][-1], reduce(lambda x, y: x * y, r[1])) for r in paths]
436 for r in paths:
437 path = [("self_loop", "user", r[0][0])]
438 for j in range(len(r[-1])):
439 path.append((r[-1][j], r[2][j], r[0][j + 1]))
440 # stop when a path is created
441 if j == len(r[-1]) - 1:
442 continue
443 pred_paths_instances[r[0][0]][r[0][-1]] = (reduce(lambda x, y: x * y, r[1]), np.mean(r[1][-1]), path)
445 top_items_scores = sorted(tmp, key=lambda x: x[1], reverse=True)
446 for item, score in top_items_scores:
447 if item < self.n_items and results[i, item] < score: # if it's an item
448 results[i, item] = score.tolist()
449 collect_results.append([user, item, score, pred_paths_instances[user][item][2]])
451 return results, collect_results
453 def create_heuristic_program(self, metapaths, predicted_paths, path_counts):
454 pcount = path_counts.astype(np.float32)
455 pcount[pcount > 5] = 5 # noqa: PLR2004
457 mp_scores = np.ones(len(metapaths)) * -99
458 for mpid in predicted_paths:
459 paths = predicted_paths[mpid]
460 if len(paths) <= 0:
461 continue
462 scores = np.array([p2[-1] for _, p2 in paths])
463 scores[scores < -5.0] = -5.0 # noqa: PLR2004
464 mp_scores[mpid] = np.mean(scores)
465 top_idxs = np.argsort(mp_scores)[::-1]
467 norm_count = np.zeros(len(metapaths))
468 rest = self.sample_size
469 for mpid in top_idxs:
470 if pcount[mpid] <= rest:
471 norm_count[mpid] = pcount[mpid]
472 else:
473 norm_count[mpid] = rest
474 rest -= norm_count[mpid]
476 program_layout = NeuralProgramLayout(metapaths)
477 program_layout.update_by_path_count(norm_count)
479 return program_layout
482class SymbolicNetwork(nn.Module):
483 def __init__(
484 self, relation_info, relation2rid, embeddings, embedding_size, deep_module, use_dropout, n_items, device
485 ):
486 super().__init__()
487 self.embedding = embeddings
488 self.embedding_size = embedding_size
489 self.n_items = n_items
490 self.device = device
491 self.relation2rid = relation2rid
493 self._create_modules(relation_info, deep_module, use_dropout)
494 self.ce_loss = nn.CrossEntropyLoss()
496 def _create_modules(self, relation_info, use_deep=False, use_dropout=True):
497 """Create module for each relation."""
498 for name in relation_info:
499 info = relation_info[name]
500 if not use_deep:
501 module = RelationModule(self.embedding_size, info)
502 else:
503 module = DeepRelationModule(self.embedding_size, info, use_dropout)
504 setattr(self, name, module)
506 def _get_modules(self, metapath):
507 """Get list of modules by metapath."""
508 module_seq = [] # seq len = len(metapath)-1
509 for relation, _ in metapath[1:]:
510 module = getattr(self, relation)
511 module_seq.append(module)
512 return module_seq
514 def _forward(self, modules, uids):
515 outputs = []
516 batch_size = uids.size(0)
518 user_vec = self.embedding["user"][uids] # [bs, d]
519 input_vec = user_vec
520 for module in modules:
521 out = module((input_vec, user_vec)).view(batch_size, -1) # [bs, d]
522 outputs.append(out)
523 input_vec = out
524 return outputs
526 def forward(self, metapath, pos_paths, neg_pids):
527 """Compute loss.
529 Args:
530 metapath: list of relations, e.g. [USER, (r1, e1),..., (r_n, e_n)].
531 pos_paths: a LongTensor of node ids, with size [bs, len(metapath)],
532 e.g. each path contains [u, e1,..., e_n].
533 neg_pids: a LongTensor of negative product ids.
535 Returns:
536 logprobs: sum of log probabilities of given target node ids, with size [bs, ].
537 """
538 modules = self._get_modules(metapath)
539 outputs = self._forward(modules, pos_paths[:, 0])
541 # Path regularization loss
542 reg_loss = 0
543 scores = 0
544 for i, module in enumerate(modules):
545 et_vecs = self.embedding[module.et_name]
546 scores = torch.matmul(outputs[i], et_vecs.t())
547 reg_loss += self.ce_loss(scores, pos_paths[:, i + 1])
549 # Ranking loss
550 logprobs = F.log_softmax(scores, dim=1) # [bs, vocab_size]
551 pos_score = torch.gather(logprobs, 1, pos_paths[:, -1].view(-1, 1))
552 neg_score = torch.gather(logprobs, 1, neg_pids.view(-1, 1))
553 rank_loss = torch.sigmoid(neg_score - pos_score).mean()
555 return reg_loss, rank_loss
557 def forward_simple(self, metapath, uids, pids):
558 modules = self._get_modules(metapath)
559 outputs = self._forward(modules, uids)
561 # Path regularization loss
562 items = self.embedding["entity"].weight[: self.n_items] # [bs, d]
563 scores = torch.matmul(outputs[-1], items.t()) # [bs, vocab_size]
564 logprobs = F.log_softmax(scores, dim=1) # [bs, vocab_size]
565 pid_logprobs = logprobs.gather(1, pids.view(-1, 1)).view(-1)
566 return pid_logprobs
568 def infer_direct(self, metapath, uid, pids):
569 if len(pids) == 0:
570 return []
571 modules = self._get_modules(metapath)
572 uid_tensor = torch.LongTensor([uid]).to(self.device)
573 # list of tensor of [1, d]
574 outputs = self._forward(modules, uid_tensor)
576 # Path regularization loss
577 pids_tensor = torch.LongTensor(pids).to(self.device)
578 items = self.embedding["entity"].weight[: self.n_items] # [bs, d]
579 scores = torch.matmul(outputs[-1], items.t()) # [1, vocab_size]
580 logprobs = F.log_softmax(scores, dim=1) # [1, vocab_size]
581 pid_logprobs = logprobs[0][pids_tensor]
582 x = pid_logprobs.detach().cpu().numpy().tolist()
583 del uid_tensor
584 del pids_tensor
585 return x
587 def infer_with_path(self, metapath, uid, kg_mask, excluded_pids, topk_paths):
588 """Reasoning paths over kg."""
589 modules = self._get_modules(metapath)
590 uid_tensor = torch.LongTensor([uid]).to(self.device)
591 # list of tensor of [1, d]
592 outputs = self._forward(modules, uid_tensor)
594 layer_logprobs = []
595 for i, module in enumerate(modules):
596 et_vecs = self.embedding[module.et_name]
597 scores = torch.matmul(outputs[i], et_vecs.t()) # [1, vocab_size]
598 logprobs = F.log_softmax(scores[0], dim=0) # [vocab_size, ]
599 layer_logprobs.append(logprobs)
601 # Decide adaptive sampling size.
602 num_valid_ids = len(kg_mask.get_ids("user", uid, self.relation2rid[modules[0].name]))
603 if num_valid_ids <= 0:
604 return []
605 sample_sizes = [topk_paths, 5, 1]
607 result_paths = [([uid], [])] # (list of ids, list of scores)
608 for i, module in enumerate(modules): # iterate over each level
609 tmp_paths = []
610 visited_ids = []
611 for path, value in result_paths: # both are lists
612 # Find valid node ids that are unvisited and not excluded pids.
613 valid_et_ids = kg_mask.get_ids(module.eh_name, path[-1], self.relation2rid[module.name])
614 valid_et_ids = set(valid_et_ids).difference(visited_ids)
615 if i == len(modules) - 1 and excluded_pids is not None:
616 valid_et_ids = valid_et_ids.difference(excluded_pids)
617 if len(valid_et_ids) <= 0:
618 continue
619 valid_et_ids = list(valid_et_ids)
621 # Compute top k nodes.
622 valid_et_ids = torch.LongTensor(valid_et_ids).to(self.device)
623 valid_et_logprobs = layer_logprobs[i].index_select(0, valid_et_ids)
624 k = min(sample_sizes[i], len(valid_et_ids))
625 topk_et_logprobs, topk_idxs = valid_et_logprobs.topk(k)
626 topk_et_ids = valid_et_ids.index_select(0, topk_idxs)
628 # Add nodes to path separately.
629 topk_et_ids = topk_et_ids.detach().cpu().numpy()
630 topk_et_logprobs = topk_et_logprobs.detach().cpu().numpy()
631 for j in range(topk_et_ids.shape[0]):
632 new_path = path + [topk_et_ids[j]]
633 new_value = value + [topk_et_logprobs[j]]
634 tmp_paths.append((new_path, new_value))
635 # Remember to add the node to visited list!!!
636 visited_ids.append(topk_et_ids[j])
637 del valid_et_ids
638 if len(tmp_paths) <= 0:
639 return []
640 result_paths = tmp_paths
641 del uid_tensor
642 return result_paths
645class RelationModule(nn.Module):
646 def __init__(self, embedding_size, relation_info):
647 super().__init__()
648 self.name = relation_info["name"]
649 self.eh_name = relation_info["entity_head"]
650 self.et_name = relation_info["entity_tail"]
651 self.fc1 = nn.Linear(embedding_size * 2, 256)
652 self.bn1 = nn.BatchNorm1d(256)
653 self.fc2 = nn.Linear(256, embedding_size)
654 self.dropout = nn.Dropout(0.5)
656 def forward(self, inputs):
657 """Compute log probability of output entity.
658 Args:
659 x: a FloatTensor of size [bs, input_size].
660 Returns:
661 FloatTensor of log probability of size [bs, output_size].
662 """
663 eh_vec, user_vec = inputs
664 x = torch.cat([eh_vec, user_vec], dim=-1)
665 x = self.bn1(self.dropout(F.relu(self.fc1(x))))
666 out = self.fc2(x) + eh_vec
667 return out
670class DeepRelationModule(nn.Module):
671 def __init__(self, embedding_size, relation_info, use_dropout):
672 super().__init__()
673 self.name = relation_info["name"]
674 self.eh_name = relation_info["entity_head"]
675 self.et_name = relation_info["entity_tail"]
676 input_size = embedding_size * 2
677 self.fc1 = nn.Linear(input_size, 256)
678 self.bn1 = nn.BatchNorm1d(256)
679 self.fc2 = nn.Linear(256, input_size)
680 self.bn2 = nn.BatchNorm1d(input_size)
681 self.fc3 = nn.Linear(input_size, embedding_size)
682 if use_dropout:
683 self.dropout = nn.Dropout(0.5)
684 else:
685 self.dropout = None
687 def forward(self, inputs):
688 eh_vec, user_vec = inputs
689 feature = torch.cat([eh_vec, user_vec], dim=-1)
690 x = F.relu(self.fc1(feature))
691 if self.dropout is not None:
692 x = self.dropout(x)
693 x = self.bn1(x)
694 x = F.relu(self.fc2(x) + feature)
695 if self.dropout is not None:
696 x = self.dropout(x)
697 x = self.bn2(x)
698 out = self.fc3(x)
699 return out
702class ReplayMemory:
703 def __init__(self, memory_size=5000):
704 self.memory_size = memory_size
705 self.memory = []
707 def add(self, data):
708 # `data` is a list of objects.
709 self.memory.extend(data)
710 while len(self.memory) > self.memory_size:
711 self.memory.pop(0)
713 def sample(self):
714 # memory is empty.
715 if not self.memory:
716 return None
717 return random.choice(self.memory)
719 def __len__(self):
720 return len(self.memory)
723class KGMask:
724 def __init__(self, kg, ui_relation_id):
725 self.kg = kg
726 self.ui_relation_id = ui_relation_id
728 def _get_next_node_type(self, current_node_type, relation_id):
729 if current_node_type == "entity" and relation_id == self.ui_relation_id:
730 return "user"
731 else:
732 return "entity"
734 def get_ids(self, eh, eh_ids, relation):
735 et_ids = []
736 if isinstance(eh_ids, list):
737 for eh_id in eh_ids:
738 try:
739 ids = list(self.kg[eh][eh_id][relation])
740 except KeyError:
741 ids = []
742 et_ids.extend(ids)
743 et_ids = list(set(et_ids))
744 else:
745 try:
746 res = self.kg[eh][eh_ids][relation]
747 except KeyError:
748 res = []
749 et_ids = list(res)
750 return et_ids
752 def get_mask(self, eh, eh_ids, relation):
753 et = self._get_next_node_type(eh, relation)
754 et_vocab_size = len(self.kg[et])
756 if isinstance(eh_ids, list):
757 mask = np.zeros([len(eh_ids), et_vocab_size], dtype=np.int64)
758 for i, eh_id in enumerate(eh_ids):
759 try:
760 et_ids = list(self.kg[eh][eh_id][relation])
761 except KeyError:
762 et_ids = []
763 mask[i, et_ids] = 1
764 else:
765 mask = np.zeros(et_vocab_size, dtype=np.int64)
766 try:
767 et_ids = list(self.kg[eh][eh_ids][relation])
768 except KeyError:
769 et_ids = []
770 mask[et_ids] = 1
771 return mask
773 def __call__(self, eh, eh_ids, relation):
774 return self.get_mask(eh, eh_ids, relation)
777class MetaProgramExecutor:
778 """This implements the profile-guided reasoning algorithm."""
780 def __init__(self, symbolic_model, random_generator, device, kg_mask, relation2rid):
781 self.symbolic_model = symbolic_model
782 self.kg_mask = kg_mask
783 self.device = device
784 self.relation2rid = relation2rid
785 self.rng = random_generator
787 def _get_module(self, relation):
788 return getattr(self.symbolic_model, relation)
790 def execute(self, program, uid, excluded_pids=None, adaptive_topk=False, manual_topk=5):
791 """Execute the program to generate node representations and real nodes.
792 Args:
793 program: an instance of MetaProgram.
794 uid: user ID (integer).
795 excluded_pids: list of item IDs (list).
796 """
797 uid_tensor = torch.LongTensor([uid]).to(self.device)
798 user_vec = self.symbolic_model.embedding["user"][uid_tensor] # tensor [1, d]
799 root = program.root # TreeNode
800 root.data["vec"] = user_vec # tensor [1, d]
801 root.data["paths"] = [([uid], [], [], [])] # (path, value, mp)
803 excluded_pids = [] if excluded_pids is None else excluded_pids.tolist()
805 # Run BFS to traverse tree.
806 queue = root.get_children()
807 while queue: # queue is not empty
808 node = queue.pop(0)
809 child_nodes = self.rng.permutation(node.get_children())
810 queue.extend(child_nodes)
812 # Compute estimated vector of the node.
813 x = (node.parent.data["vec"], user_vec)
814 node.data["vec"] = self._get_module(node.relation)(x) # tensor [1, d]
816 # Compute scores (log prob) for the node.
817 entity_vecs = self.symbolic_model.embedding[node.entity] # tensor [vocab, d]
818 # tensor [1, vocab]
819 scores = torch.matmul(node.data["vec"], entity_vecs.t())
820 scores = F.log_softmax(scores[0], dim=0) # tensor [vocab, ]
822 node.data["paths"] = []
823 visited_ids = []
824 for path, value, ep, mp in node.parent.data["paths"]:
825 # Find valid node ids for current path.
826 valid_ids = self.kg_mask.get_ids(node.parent.entity, path[-1], self.relation2rid[node.relation])
827 valid_ids = set(valid_ids).difference(visited_ids)
828 if not node.has_children() and excluded_pids:
829 valid_ids = valid_ids.difference(excluded_pids)
830 if not valid_ids: # empty list
831 continue
832 valid_ids = list(valid_ids)
834 # Compute top k nodes.
835 valid_ids = torch.LongTensor(valid_ids).to(self.device)
836 valid_scores = scores.index_select(0, valid_ids)
837 if adaptive_topk:
838 k = min(node.sample_size, len(valid_ids))
839 else:
840 k = min(manual_topk, len(valid_ids))
841 topk_scores, topk_idxs = valid_scores.topk(k)
842 topk_ids = valid_ids.index_select(0, topk_idxs)
844 # Add nodes and scores to paths.
845 topk_ids = topk_ids.detach().cpu().numpy()
846 topk_scores = topk_scores.detach().cpu().numpy()
847 for j in range(k):
848 new_path = path + [topk_ids[j]]
849 new_value = value + [topk_scores[j]]
850 new_mp = mp + [node.relation]
851 new_ep = ep + [node.entity]
852 node.data["paths"].append((new_path, new_value, new_ep, new_mp))
854 # Remember to add the node to visited list!!!
855 visited_ids.append(topk_ids[j])
856 if not node.has_children():
857 excluded_pids.append(topk_ids[j])
859 def collect_results(self, program):
860 results = []
861 queue = program.root.get_children()
862 while len(queue) > 0:
863 node = queue.pop(0)
864 queue.extend(node.get_children())
865 if not node.has_children():
866 results.extend(node.data["paths"])
867 return results
870class NeuralProgramLayout:
871 """This refers to the layout tree in the paper."""
873 def __init__(self, metapaths):
874 super().__init__()
875 self.mp2id = {}
876 for mpid, mp in enumerate(metapaths):
877 simple_mp = tuple([v[0] for v in mp[1:]])
878 self.mp2id[simple_mp] = mpid
880 self.root = TreeNode(0, "user", None)
881 for mp in metapaths:
882 node = self.root
883 for i in range(1, len(mp)):
884 if mp[i] not in node.children:
885 node.children[mp[i]] = TreeNode(i, mp[i][1], mp[i][0])
886 node.children[mp[i]].parent = node
887 node = node.children[mp[i]]
889 def update_by_path_count(self, path_count):
890 """Update sample size of each node by expected number of paths.
891 Args:
892 path_count: dict with key=mpid, value=int
893 """
895 def _postorder_update(node, parent_rels):
896 if not node.has_children():
897 mpid = self.mp2id[tuple(parent_rels)]
898 node.sample_size = int(path_count[mpid])
899 return
901 min_pos_sample_size, max_sample_size = 99, 0
902 for child in node.get_children():
903 _postorder_update(child, parent_rels + [child.relation])
904 max_sample_size = max(max_sample_size, child.sample_size)
905 if child.sample_size > 0:
906 min_pos_sample_size = min(min_pos_sample_size, child.sample_size)
908 # Update current node sampling size.
909 # a) if current node is root, set to 1.
910 if not node.has_parent():
911 node.sample_size = 1
912 # b) if current node is not root, and all children sample sizes are 0, set to 0.
913 elif max_sample_size == 0:
914 node.sample_size = 0
915 # c) if current node is not root, take the minimum and update children.
916 else:
917 node.sample_size = min_pos_sample_size
918 for child in node.get_children():
919 child.sample_size = int(child.sample_size / node.sample_size)
921 _postorder_update(self.root, [])
923 def print_postorder(self, hide_branch=True):
924 def _postorder(node, msgs):
925 msg = (node.entity, node.relation, node.sample_size)
926 new_msgs = msgs + [msg]
928 if not node.has_children():
929 if hide_branch and msg[2] == 0:
930 return
931 str_msgs = [f"({msg[0]},{msg[1]},{msg[2]})" for msg in new_msgs]
932 print(" ".join(str_msgs))
933 return
935 for child in node.children:
936 _postorder(child, new_msgs)
938 _postorder(self.root, [])
941class TreeNode:
942 def __init__(self, level, entity, relation):
943 super().__init__()
944 self.level = level
945 self.entity = entity # Entity type
946 self.relation = relation # Relation pointing to this tail entity
947 self.parent = None
948 self.children = {} # key = (relation, entity), value = TreeNode
949 self.sample_size = 0 # number of nodes to sample
950 self.data = {} # extra information to save
952 def has_parent(self):
953 return self.parent is not None
955 def has_children(self):
956 return len(self.children) > 0
958 def get_children(self):
959 return list(self.children.values())
961 def __str__(self):
962 parent = None if not self.has_parent() else self.parent.entity
963 msg = f"({parent},{self.relation},{self.entity})"
964 return msg