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

1# @Time : 2025/02/19 

2# @Author : Alessandro Soccol 

3# @Email : alessandro.soccol@unica.it 

4 

5"""CAFE 

6################################################## 

7Reference: 

8 Xian et al. "CAFE: Coarse-to-Fine Neural Symbolic Reasoning for Explainable Recommendation." in CIKM 2020. 

9 

10Reference code: 

11 https://github.com/orcax/CAFE 

12""" 

13 

14import random 

15from functools import reduce 

16 

17import numpy as np 

18import torch 

19import torch.nn.functional as F 

20from torch import nn 

21 

22from hopwise.model.abstract_recommender import KnowledgeRecommender 

23from hopwise.utils import InputType 

24 

25 

26class CAFE(KnowledgeRecommender): 

27 """ 

28 CAFE is a knowledge-aware recommender system that uses symbolic reasoning 

29 over a knowledge graph to explain recommendations. 

30 

31 Note: 

32 Assumes that each relation corresponds to a unique pair of entity types. e.g. ui-relation -> (user, item) 

33 """ 

34 

35 input_type = InputType.USERWISE 

36 

37 def __init__(self, config, dataset): 

38 super().__init__(config, dataset) 

39 self.dataset = dataset 

40 

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"] 

45 

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"] 

54 

55 # user-item relation 

56 self.ui_relation = dataset.ui_relation 

57 self.ui_relation_id = dataset.field2token_id["relation_id"][dataset.ui_relation] 

58 

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") 

63 

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) 

71 

72 # Embedding mapping 

73 self.embeddings = { 

74 "user": self.user_embedding, 

75 "entity": self.entity_embedding, 

76 "relation": self.relation_embedding, 

77 } 

78 

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"] 

82 

83 # Positives 

84 self.positives = dataset.history_item_matrix()[0] 

85 

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 = {} 

90 

91 self.relation_info = dict() 

92 for relation_name in self.rid2relation.values(): 

93 if relation_name == "[PAD]": 

94 continue 

95 

96 if relation_name == f"{dataset.ui_relation}_r": 

97 raise ValueError("The ui_relation name should not end with '_r'.") 

98 

99 if relation_name == dataset.ui_relation: 

100 head, tail = "user", "entity" 

101 else: 

102 head, tail = "entity", "entity" 

103 

104 if relation_name not in self.relation_info: 

105 self.relation_info[relation_name] = {"name": relation_name, "entity_head": head, "entity_tail": tail} 

106 

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) 

114 

115 self.mpath_ids = list(range(len(self.metapaths))) 

116 

117 for mpid in range(len(self.metapaths)): 

118 self.replay_memory[mpid] = ReplayMemory(self.memory_size) 

119 

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 ) 

130 

131 # random generator 

132 self.rng = np.random.default_rng() 

133 

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 

142 

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() 

153 

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) 

162 

163 # Sample one of the best topk_candidates e.g. sample one of the top20 pids 

164 pidx = self.rng.choice(top_pids) 

165 

166 # Take the corresponding score 

167 item = self.topk_user_items[user][pidx] 

168 

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 

171 

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) 

185 

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) 

192 

193 neg_pid_batch.append(neg_pid) 

194 it += 1 

195 

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 

199 

200 def _rev_rel(self, rel): 

201 if rel == self.ui_relation: 

202 return self.ui_relation 

203 

204 if rel.endswith("_r"): 

205 return rel[:-2] 

206 return rel + "_r" 

207 

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. 

210 

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) 

217 

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 

286 

287 return final_paths 

288 

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) 

296 

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 

312 

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 

325 

326 count = len(set(forward_ids).intersection(backward_ids)) 

327 return count 

328 

329 def calculate_loss(self, interaction): 

330 users = interaction[self.USER_ID] 

331 users = users[users != 0] 

332 

333 mpid, pos_paths, neg_pids = self._get_batch_by_user(users) 

334 

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 

339 

340 return reg_loss, rank_loss 

341 

342 def predict(self, interaction): 

343 return 

344 

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 

354 

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 

363 

364 def explain(self, interaction): 

365 """Support function used for case study. 

366 

367 Args: 

368 interaction : test interaction data 

369 

370 Returns: 

371 pd.Dataframe: explanation results with columns: "user", "item", "score", "path" 

372 """ 

373 users = interaction[self.USER_ID] 

374 

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) 

378 

379 scores, explanations = self.run_program(users, path_counts, predicted_paths) 

380 

381 for exp in explanations: 

382 exp[-1] = self.decode_path(exp[-1]) 

383 

384 return scores, explanations 

385 

386 def decode_path(self, path): 

387 return path 

388 

389 def _infer_paths(self, users, kg_mask): 

390 predictions = dict() 

391 for user in users: 

392 user = user.item() # noqa: PLW2901 

393 

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 

401 

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 

407 

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 

417 

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() 

421 

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) 

424 

425 pred_paths_instances = dict() 

426 for i, user in enumerate(users): 

427 user = user.item() # noqa: PLW2901 

428 

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) 

444 

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]]) 

450 

451 return results, collect_results 

452 

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 

456 

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] 

466 

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] 

475 

476 program_layout = NeuralProgramLayout(metapaths) 

477 program_layout.update_by_path_count(norm_count) 

478 

479 return program_layout 

480 

481 

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 

492 

493 self._create_modules(relation_info, deep_module, use_dropout) 

494 self.ce_loss = nn.CrossEntropyLoss() 

495 

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) 

505 

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 

513 

514 def _forward(self, modules, uids): 

515 outputs = [] 

516 batch_size = uids.size(0) 

517 

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 

525 

526 def forward(self, metapath, pos_paths, neg_pids): 

527 """Compute loss. 

528 

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. 

534 

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]) 

540 

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]) 

548 

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() 

554 

555 return reg_loss, rank_loss 

556 

557 def forward_simple(self, metapath, uids, pids): 

558 modules = self._get_modules(metapath) 

559 outputs = self._forward(modules, uids) 

560 

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 

567 

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) 

575 

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 

586 

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) 

593 

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) 

600 

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] 

606 

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) 

620 

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) 

627 

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 

643 

644 

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) 

655 

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 

668 

669 

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 

686 

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 

700 

701 

702class ReplayMemory: 

703 def __init__(self, memory_size=5000): 

704 self.memory_size = memory_size 

705 self.memory = [] 

706 

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) 

712 

713 def sample(self): 

714 # memory is empty. 

715 if not self.memory: 

716 return None 

717 return random.choice(self.memory) 

718 

719 def __len__(self): 

720 return len(self.memory) 

721 

722 

723class KGMask: 

724 def __init__(self, kg, ui_relation_id): 

725 self.kg = kg 

726 self.ui_relation_id = ui_relation_id 

727 

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" 

733 

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 

751 

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]) 

755 

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 

772 

773 def __call__(self, eh, eh_ids, relation): 

774 return self.get_mask(eh, eh_ids, relation) 

775 

776 

777class MetaProgramExecutor: 

778 """This implements the profile-guided reasoning algorithm.""" 

779 

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 

786 

787 def _get_module(self, relation): 

788 return getattr(self.symbolic_model, relation) 

789 

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) 

802 

803 excluded_pids = [] if excluded_pids is None else excluded_pids.tolist() 

804 

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) 

811 

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] 

815 

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, ] 

821 

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) 

833 

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) 

843 

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)) 

853 

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]) 

858 

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 

868 

869 

870class NeuralProgramLayout: 

871 """This refers to the layout tree in the paper.""" 

872 

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 

879 

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]] 

888 

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 """ 

894 

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 

900 

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) 

907 

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) 

920 

921 _postorder_update(self.root, []) 

922 

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] 

927 

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 

934 

935 for child in node.children: 

936 _postorder(child, new_msgs) 

937 

938 _postorder(self.root, []) 

939 

940 

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 

951 

952 def has_parent(self): 

953 return self.parent is not None 

954 

955 def has_children(self): 

956 return len(self.children) > 0 

957 

958 def get_children(self): 

959 return list(self.children.values()) 

960 

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