Coverage for hopwise/model/knowledge_aware_recommender/tprec.py: 9%

522 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1# @Time : 2025/05/28 

2# @Author : Alessandro Soccol 

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

4 

5"""TPRec 

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

7Reference: Time-aware Path Reasoning on Knowledge Graph for Recommendation (https://arxiv.org/pdf/2108.02634) 

8 

9""" 

10 

11from collections import defaultdict, namedtuple 

12from functools import reduce 

13 

14import numpy as np 

15import torch 

16import torch.nn.functional as F 

17from torch import nn 

18from torch.distributions import Categorical 

19 

20from hopwise.model.abstract_recommender import ExplainableRecommender, KnowledgeRecommender 

21from hopwise.utils import InputType 

22 

23 

24class TPRec(KnowledgeRecommender, ExplainableRecommender): 

25 """ 

26 TPRec 

27 

28 1. Train TransE embeddings and preprocess it according to the preprocess embedding notebook 

29 2. Run TPRec with 'pretrain' train stage set 

30 3. Run TPRec with 'policy' train stage set 

31 

32 """ 

33 

34 input_type = InputType.USERWISE 

35 

36 def __init__(self, config, dataset): 

37 super().__init__(config, dataset) 

38 self.config = config 

39 

40 # Load parameters info from config 

41 self.user_num = dataset.user_num 

42 self.device = config["device"] 

43 self.topk = config["topk"] 

44 

45 # PGPR Configurations 

46 self.state_history = config["state_history"] 

47 self.max_acts = config["max_acts"] 

48 self.gamma = config["gamma"] 

49 self.action_dropout = config["action_dropout"] 

50 self.hidden_sizes = config["hidden_sizes"] 

51 self.act_dim = self.max_acts + 1 

52 self.max_num_nodes = config["max_path_len"] + 1 

53 self.weight_factor = config["weight_factor"] 

54 self.path_pattern = config["path_constraint"] 

55 self.beam_search_hop = config["beam_search_hop"] 

56 self.train_stage = config["train_stage"] 

57 self.margin = config["margin"] 

58 self.n_clusters = config["cluster_num"] 

59 self.fix_scores_sorting_bug = config["fix_scores_sorting_bug"] 

60 

61 # user-item relation 

62 self.ui_relation = dataset.ui_relation 

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

64 

65 self.graph_dict = dataset.ckg_dict_graph() 

66 # Items 

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

68 self.pretrained_weights = {} # during pretraining is None, otherwise it contain pretrained weights 

69 self.uc_weight = dataset.temporal_weights.uc_weight 

70 self.timenum = dataset.temporal_weights.timenum 

71 self.ui2label_dict = dataset.temporal_weights.timeClassifyLabel 

72 

73 # Load Knowledge Graph Embedding Checkpoint 

74 if self.train_stage == "pretrain": 

75 # then load transe embeddings and initialize new torch embeddings 

76 self.user_embedding = dataset.get_preload_weight("user_embedding_id") 

77 self.entity_embedding = dataset.get_preload_weight("entity_embedding_id") 

78 self.relation_embedding = dataset.get_preload_weight("relation_embedding_id") 

79 

80 # make embeddings learnable 

81 self.user_embedding = torch.from_numpy(self.user_embedding) 

82 self.entity_embedding = torch.from_numpy(self.entity_embedding) 

83 self.relation_embedding = torch.from_numpy(self.relation_embedding) 

84 

85 self.user_embedding = nn.Embedding.from_pretrained(self.user_embedding) 

86 self.entity_embedding = nn.Embedding.from_pretrained(self.entity_embedding) 

87 self.relation_embedding = nn.Embedding.from_pretrained(self.relation_embedding) 

88 

89 # TPRec temporal embeddings 

90 embedding_size = self.user_embedding.weight.size(1) 

91 self.ui_clust_relation_embedding = nn.Embedding(self.n_clusters, embedding_size) 

92 

93 self.transe_loss = nn.TripletMarginLoss(margin=config["margin"], p=2, reduction="mean") 

94 return 

95 else: 

96 # load pretrained weights 

97 pretrained_weights = self._get_pretrained_weights() 

98 # then transform torch embedding in numpy 

99 # transform torch.load self.pretrained_weights in numpy 

100 self.user_embedding = pretrained_weights["user_embedding.weight"].cpu().numpy() 

101 self.entity_embedding = pretrained_weights["entity_embedding.weight"].cpu().numpy() 

102 self.relation_embedding = pretrained_weights["relation_embedding.weight"].cpu().numpy() 

103 self.ui_clust_relation_embedding = pretrained_weights["ui_clust_relation_embedding.weight"].cpu().numpy() 

104 self.embedding_size = self.user_embedding.shape[1] 

105 self.state_gen = KGState(self.embedding_size, self.state_history) 

106 

107 # Actor-Critic model 

108 self.l1 = nn.Linear(self.state_gen.dim, self.hidden_sizes[0]) 

109 self.l2 = nn.Linear(self.hidden_sizes[0], self.hidden_sizes[1]) 

110 self.actor = nn.Linear(self.hidden_sizes[1], self.act_dim) 

111 self.critic = nn.Linear(self.hidden_sizes[1], 1) 

112 

113 # Self Loop Embedding 

114 self.self_loop_embedding = np.zeros(self.embedding_size) 

115 

116 # Map Relation ID to relation name to check has pattern constraint 

117 self.rid2relation = {v: k for k, v in dataset.field2token_id["relation_id"].items()} 

118 

119 # Mapping node type to embeddings 

120 self.node_type2emb = { 

121 "user": self.user_embedding, 

122 "entity": self.entity_embedding, 

123 "relation": self.relation_embedding, 

124 "self_loop": self.self_loop_embedding, 

125 } 

126 

127 # Normalization score 

128 u_p_scores = np.dot( 

129 self.user_embedding + self.relation_embedding[self.ui_relation_id], self.entity_embedding[: self.n_items].T 

130 ) 

131 self.u_p_scales = np.max(u_p_scores, axis=1) 

132 

133 # These are the paths constraint to use when checking for path correctness in _get_reward function 

134 # Preprocess the path constraints. 

135 self.patterns = list() 

136 for path_constraint in self.path_pattern: 

137 relations = list() 

138 for node in path_constraint: 

139 path_rel = node[0] 

140 if path_rel is not None: 

141 if path_rel.endswith("_r"): 

142 # remove the reverse suffix 

143 path_rel = path_rel[:-2] 

144 relations.append(path_rel) 

145 self.patterns.append(tuple(["self_loop"]) + tuple(relations)) 

146 

147 # this second step is added separately for clarity 

148 # expand UI-Relation to the number of clusters according to path preprocessing in TPRec 

149 extended_paths = [] 

150 for path_constraint in self.patterns.copy(): 

151 extended_paths.extend(self.expand_paths(path_constraint)) 

152 self.patterns = extended_paths # remove duplicates 

153 

154 # Following is current episode information. 

155 self._batch_path = None 

156 self._batch_curr_actions = None 

157 self._batch_curr_state = None 

158 self._batch_curr_reward = None 

159 self._done = False 

160 self.saved_actions = [] 

161 self.rewards = [] 

162 self.entropy = [] 

163 

164 # random generator 

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

166 self.SavedAction = namedtuple("SavedAction", ["log_prob", "value"]) 

167 

168 def expand_paths(self, path_constraint): 

169 """ 

170 Expand the path constraint by replacing the ui_relation with multiple clusters 

171 like ui_relation0 to ui_relation_n with n being the number of clusters. 

172 

173 Args: 

174 path_constraint in the form of [rel1, rel2, rel3] 

175 """ 

176 

177 def _expand_recursive(current_path, remaining_relations): 

178 # Base case - path complete 

179 if not remaining_relations: 

180 return [current_path] 

181 

182 # Get next relation 

183 relation = remaining_relations[0] 

184 

185 # If ui_relation, branch into n_clusters paths 

186 if relation == self.ui_relation: 

187 expanded = [] 

188 for i in range(self.n_clusters): 

189 cluster_path = current_path + [f"{relation}_{i}"] 

190 expanded.extend(_expand_recursive(cluster_path, remaining_relations[1:])) 

191 return expanded 

192 else: 

193 # Regular relation - add and continue 

194 return _expand_recursive(current_path + [relation], remaining_relations[1:]) 

195 

196 # Start expansion from empty path 

197 return _expand_recursive([], path_constraint[1:]) # Skip self-loop 

198 

199 def _get_pretrained_weights(self): 

200 import os 

201 

202 checkpoint_file = os.path.join( 

203 self.config["checkpoint_dir"], 

204 "{}-{}-{}.pth".format(self.config["model"], self.config["dataset"], "pretrained"), 

205 ) 

206 checkpoint = torch.load(checkpoint_file, weights_only=False, map_location=self.device) 

207 weights = checkpoint["state_dict"] 

208 return weights 

209 

210 def select_action(self, batch_state, batch_act_mask): 

211 # Tensor [bs, state_dim] 

212 state = torch.FloatTensor(batch_state).to(self.device) 

213 # Tensor of [bs, act_dim] 

214 act_mask = torch.BoolTensor(batch_act_mask).to(self.device) 

215 # act_probs: [bs, act_dim], state_value: [bs, 1] 

216 probs, value = self.forward((state, act_mask)) 

217 m = Categorical(probs) 

218 acts = m.sample() # Tensor of [bs, ], requires_grad=False 

219 # [CAVEAT] If sampled action is out of action_space, choose the first action in action_space. 

220 valid_idx = act_mask.gather(1, acts.view(-1, 1)).view(-1) 

221 acts[valid_idx == 0] = 0 

222 

223 self.saved_actions.append(self.SavedAction(m.log_prob(acts), value)) 

224 self.entropy.append(m.entropy()) 

225 return acts.cpu().numpy().tolist() 

226 

227 def update(self): # prev update 

228 if len(self.rewards) <= 0: 

229 del self.rewards[:] 

230 del self.saved_actions[:] 

231 del self.entropy[:] 

232 return 0.0, 0.0, 0.0 

233 

234 # numpy array of [bs, #steps] 

235 batch_rewards = np.vstack(self.rewards).T 

236 batch_rewards = torch.tensor(batch_rewards).to(self.device) 

237 num_steps = batch_rewards.shape[1] 

238 

239 for i in range(1, num_steps): 

240 batch_rewards[:, num_steps - i - 1] += self.gamma * batch_rewards[:, num_steps - i] 

241 

242 actor_loss = 0 

243 critic_loss = 0 

244 entropy_loss = 0 

245 

246 for i in range(0, num_steps): 

247 # log_prob: Tensor of [bs, ], value: Tensor of [bs, 1] 

248 log_prob, value = self.saved_actions[i] 

249 advantage = batch_rewards[:, i] - value.squeeze(1) # Tensor of [bs, ] 

250 actor_loss += -log_prob * advantage.detach() # Tensor of [bs, ] 

251 critic_loss += advantage.pow(2) # Tensor of [bs, ] 

252 entropy_loss += -self.entropy[i] # Tensor of [bs, ] 

253 

254 actor_loss = actor_loss.mean() 

255 critic_loss = critic_loss.mean() 

256 entropy_loss = entropy_loss.mean() 

257 loss = actor_loss + critic_loss + self.weight_factor * entropy_loss 

258 

259 del self.rewards[:] 

260 del self.saved_actions[:] 

261 del self.entropy[:] 

262 

263 return loss, actor_loss, critic_loss, entropy_loss 

264 

265 def forward(self, inputs): 

266 # used only in inference 

267 # state: [bs, state_dim], act_mask: [bs, act_dim] 

268 state, act_mask = inputs 

269 x = self.l1(state) 

270 x = F.dropout(F.elu(x), p=0.5) 

271 out = self.l2(x) 

272 x = F.dropout(F.elu(out), p=0.5) 

273 actor_logits = self.actor(x) 

274 actor_logits[~act_mask] = float("-inf") 

275 act_probs = F.softmax(actor_logits, dim=-1) # Tensor of [bs, act_dim] 

276 state_values = self.critic(x) # Tensor of [bs, 1] 

277 return act_probs, state_values 

278 

279 def _get_transe_rec_embedding(self, user, pos_item, neg_item): 

280 user_e = self.user_embedding(user) 

281 pos_item_e = self.entity_embedding(pos_item) 

282 neg_item_e = self.entity_embedding(neg_item) 

283 rec_r_e = self.relation_embedding.weight[-1] 

284 rec_r_e = rec_r_e.expand_as(user_e) 

285 

286 return user_e, pos_item_e, neg_item_e, rec_r_e 

287 

288 def _get_transe_kg_embedding(self, head, pos_tail, neg_tail, relation): 

289 head_e = self.entity_embedding(head) 

290 pos_tail_e = self.entity_embedding(pos_tail) 

291 neg_tail_e = self.entity_embedding(neg_tail) 

292 relation_e = self.relation_embedding(relation) 

293 return head_e, pos_tail_e, neg_tail_e, relation_e 

294 

295 def calculate_loss_transe(self, interaction): 

296 user = interaction[self.USER_ID] 

297 pos_item = interaction[self.ITEM_ID] 

298 ui_clusters_nums = [ 

299 self.ui2label_dict[(user.item(), pos_item.item())] for user, pos_item in zip(user, pos_item) 

300 ] 

301 num_id, num_index, num_n = np.unique(ui_clusters_nums, return_index=True, return_counts=True) 

302 neg_item = interaction[self.NEG_ITEM_ID] 

303 head = interaction[self.HEAD_ENTITY_ID] 

304 relation = interaction[self.RELATION_ID] 

305 pos_tail = interaction[self.TAIL_ENTITY_ID] 

306 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID] 

307 

308 user_e, pos_item_e, neg_item_e, rec_r_e = self._get_transe_rec_embedding(user, pos_item, neg_item) 

309 head_e, pos_tail_e, neg_tail_e, relation_e = self._get_transe_kg_embedding(head, pos_tail, neg_tail, relation) 

310 

311 loss = 0 

312 if len(num_id) == 1: 

313 rec_r_e = self.ui_clust_relation_embedding.weight[num_id[0]] 

314 else: 

315 tmp_vec_rel = None 

316 for i, cur_cls in enumerate(num_id): 

317 startID = num_index[i] 

318 endID = num_index[i] + num_n[i] 

319 rec_clus_r_e = self.ui_clust_relation_embedding.weight[cur_cls] 

320 if tmp_vec_rel is None: 

321 tmp_vec_rel = rec_r_e 

322 else: 

323 torch.cat((rec_r_e[0], tmp_vec_rel[0]), 0) 

324 loss += self.transe_loss( 

325 user_e[startID:endID] + rec_clus_r_e, pos_item_e[startID:endID], neg_item_e[startID:endID] 

326 ) 

327 h_e = torch.cat([user_e, head_e]) 

328 r_e = torch.cat([rec_r_e, relation_e]) 

329 pos_t_e = torch.cat([pos_item_e, pos_tail_e]) 

330 neg_t_e = torch.cat([neg_item_e, neg_tail_e]) 

331 

332 loss += self.transe_loss(h_e + r_e, pos_t_e, neg_t_e) 

333 

334 return loss 

335 

336 def forward_transe(self, user, relation, item): 

337 score = -torch.norm(user + relation - item, p=2, dim=1) 

338 return score 

339 

340 def predict_transe(self, interaction): 

341 user = interaction[self.USER_ID] 

342 item = interaction[self.ITEM_ID] 

343 

344 user_e = self.user_embedding(user) 

345 item_e = self.entity_embedding(item) 

346 

347 rec_r_e = self.relation_embedding.weight[-1] 

348 rec_r_e = rec_r_e.expand_as(user_e) 

349 

350 return self.forward_transe(user_e, rec_r_e, item_e) 

351 

352 def full_sort_predict_transe(self, interaction): 

353 user = interaction[self.USER_ID] 

354 user_e = self.user_embedding(user) 

355 

356 rec_r_e = self.relation_embedding.weight[-1] 

357 rec_r_e = rec_r_e.expand_as(user_e) 

358 

359 item_indices = torch.tensor(range(self.n_items)).to(self.device) 

360 all_item_e = self.entity_embedding.weight[item_indices] 

361 

362 user_e = user_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1) 

363 rec_r_e = rec_r_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1) 

364 t = all_item_e.unsqueeze(0) 

365 

366 return -torch.norm(user_e + rec_r_e - t, p=2, dim=2) 

367 

368 def calculate_loss(self, interaction): 

369 if self.train_stage == "pretrain": 

370 return self.calculate_loss_transe(interaction) 

371 

372 users = interaction[self.USER_ID] 

373 users = users[users != 0] 

374 

375 # Policy training 

376 batch_state = self.reset(users) 

377 done = False 

378 while not done: 

379 batch_act_mask = self.batch_action_mask(dropout=self.action_dropout) 

380 batch_act_idx = self.select_action(batch_state, batch_act_mask) 

381 batch_state, batch_reward, done = self.batch_step(batch_act_idx) 

382 self.rewards.append(batch_reward) 

383 loss, ploss, vloss, eloss = self.update() 

384 

385 return loss 

386 

387 def _has_pattern(self, path): 

388 pattern = tuple([self.rid2relation[v[0]] if v[0] != "self_loop" else v[0] for v in path]) 

389 return pattern in self.patterns 

390 

391 def _get_next_node_type(self, current_node_type, relation_id): 

392 if current_node_type == "entity" and relation_id == self.ui_relation_id: 

393 return "user" 

394 else: 

395 return "entity" 

396 

397 def reset(self, user): 

398 self._batch_path = [[("self_loop", "user", uid)] for uid in user] 

399 self._done = False 

400 self._batch_curr_state = self._batch_get_state(self._batch_path) 

401 self._batch_curr_actions = self._batch_get_actions(self._batch_path, self._done) 

402 self._batch_curr_reward = self._batch_get_reward(self._batch_path) 

403 return self._batch_curr_state 

404 

405 def _batch_get_actions(self, batch_path, done): 

406 return [self._get_actions(path, done) for path in batch_path] 

407 

408 def _get_actions(self, path, done): 

409 # Compute actions for current node. 

410 _, curr_node_type, curr_node_id = path[-1] 

411 actions = [("self_loop", curr_node_id)] 

412 

413 # (1) If game is finished, only return self-loop action. 

414 if done: 

415 return actions 

416 

417 # (2) Get all possible edges from original knowledge graph. 

418 # [CAVEAT] Must remove visited nodes! 

419 if isinstance(curr_node_id, torch.Tensor): 

420 curr_node_id = curr_node_id.item() 

421 

422 try: 

423 relations_nodes = self.graph_dict[curr_node_type][curr_node_id] 

424 except KeyError: 

425 relations_nodes = [] 

426 candidate_acts = [] # list of tuples of (relation, node_type, node_id) 

427 visited_nodes = set([(v[1], v[2]) for v in path]) 

428 

429 for r in relations_nodes: 

430 next_node_type = self._get_next_node_type(curr_node_type, r) 

431 next_node_ids = relations_nodes[r] 

432 next_node_ids = [n for n in next_node_ids if (next_node_type, n) not in visited_nodes] # filter 

433 candidate_acts.extend(zip([r] * len(next_node_ids), next_node_ids)) 

434 

435 # (3) If candidate action set is empty, only return self-loop action. 

436 if len(candidate_acts) == 0: 

437 return actions 

438 

439 # (4) If number of available actions is smaller than max_acts, return action sets. 

440 if len(candidate_acts) <= self.max_acts: 

441 candidate_acts = sorted(candidate_acts, key=lambda x: (x[0], x[1])) 

442 actions.extend(candidate_acts) 

443 return actions 

444 

445 # (5) If there are too many actions, do some deterministic trimming here! 

446 uid = path[0][-1] 

447 user_embed = self.user_embedding[uid] 

448 

449 scores = [] 

450 item_emb = None 

451 if isinstance(uid, torch.Tensor): 

452 uid = uid.item() 

453 

454 for clus_wt in self.uc_weight[uid]: 

455 if item_emb is None: 

456 item_emb = self.ui_clust_relation_embedding[clus_wt] * self.uc_weight[uid][clus_wt] 

457 else: 

458 item_emb += self.ui_clust_relation_embedding[clus_wt] * self.uc_weight[uid][clus_wt] 

459 

460 for r, next_node_id in candidate_acts: 

461 next_node_type = self._get_next_node_type(curr_node_type, r) 

462 if next_node_type == "user": 

463 src_embed = user_embed 

464 elif next_node_type == "entity" and next_node_id < self.n_items: 

465 src_embed = user_embed + item_emb 

466 else: 

467 src_embed = user_embed + item_emb + self.relation_embedding[r] 

468 

469 score = np.matmul(src_embed, self.node_type2emb[next_node_type][next_node_id]) 

470 # This trimming may filter out target items! 

471 # Manually set the score of target items a very large number. 

472 # if next_node_type == ITEM and next_node_id in self._target_pids: 

473 # score = 99999.0 

474 scores.append(score) 

475 

476 # choose actions with larger scores 

477 candidate_idxs = np.argsort(scores)[-self.max_acts :] 

478 if self.fix_scores_sorting_bug: 

479 candidate_acts = [candidate_acts[i] for i in candidate_idxs[::-1]] 

480 else: 

481 candidate_acts = sorted([candidate_acts[i] for i in candidate_idxs], key=lambda x: (x[0], x[1])) 

482 actions.extend(candidate_acts) 

483 return actions 

484 

485 def _batch_get_state(self, batch_path): 

486 batch_state = [self._get_state(path) for path in batch_path] 

487 return np.vstack(batch_state) # [bs, dim] 

488 

489 def _get_state(self, path): 

490 # Return state of torch vector: [user_embed, curr_node_embed, last_node_embed, last_relation]. 

491 user_embed = self.user_embedding[path[0][-1]] 

492 zero_embed = np.zeros(self.embedding_size) 

493 

494 if len(path) == 1: # initial state 

495 state = self.state_gen(user_embed, user_embed, zero_embed, zero_embed, zero_embed, zero_embed) 

496 return state 

497 

498 older_relation, last_node_type, last_node_id = path[-2] 

499 last_relation, curr_node_type, curr_node_id = path[-1] 

500 

501 curr_node_embed = self.node_type2emb[curr_node_type][curr_node_id] 

502 last_node_embed = self.node_type2emb[last_node_type][last_node_id] 

503 

504 if last_relation != "self_loop": 

505 last_relation_embed = self.relation_embedding[last_relation] 

506 else: 

507 last_relation_embed = self.self_loop_embedding 

508 

509 if len(path) == 2: # noqa: PLR2004 

510 state = self.state_gen( 

511 user_embed, curr_node_embed, last_node_embed, last_relation_embed, zero_embed, zero_embed 

512 ) 

513 return state 

514 

515 _, older_node_type, older_node_id = path[-3] 

516 older_node_embed = self.node_type2emb[older_node_type][older_node_id] 

517 

518 if older_relation == "self_loop": 

519 older_relation_embed = self.self_loop_embedding 

520 else: 

521 older_relation_embed = self.relation_embedding[older_relation] 

522 

523 state = self.state_gen( 

524 user_embed, curr_node_embed, last_node_embed, last_relation_embed, older_node_embed, older_relation_embed 

525 ) 

526 return state 

527 

528 def _batch_get_reward(self, batch_path): 

529 batch_reward = [self._get_reward(path) for path in batch_path] 

530 return np.array(batch_reward) 

531 

532 def _get_reward(self, path): 

533 # If it is initial state or 1-hop search, reward is 0. 

534 if len(path) <= 2: # noqa: PLR2004 

535 return 0.0 

536 

537 if not self._has_pattern(path): 

538 return 0.0 

539 

540 target_score = 0.0 

541 _, curr_node_type, curr_node_id = path[-1] 

542 

543 if curr_node_type == "entity" and curr_node_id < self.n_items: 

544 # Give soft reward for other reached items. 

545 uid = path[0][-1] 

546 item_emb = None 

547 for clus_wt in self.uc_weight[uid.item()]: 

548 if item_emb is None: 

549 item_emb = self.ui_clust_relation_embedding[clus_wt] * self.uc_weight[uid.item()][clus_wt] 

550 else: 

551 item_emb += self.ui_clust_relation_embedding[clus_wt] * self.uc_weight[uid.item()][clus_wt] 

552 

553 u_vec = self.user_embedding[uid] + item_emb 

554 p_vec = self.entity_embedding[curr_node_id] 

555 score = np.dot(u_vec, p_vec) / self.u_p_scales[uid] 

556 target_score = max(score, 0.0) 

557 return target_score 

558 

559 def _is_done(self): 

560 # Episode ends only if max path length is reached. 

561 return self._done or len(self._batch_path[0]) >= self.max_num_nodes 

562 

563 def batch_step(self, batch_act_idx): 

564 assert len(batch_act_idx) == len(self._batch_path) 

565 

566 # Execute batch actions. 

567 for i in range(len(batch_act_idx)): 

568 act_idx = batch_act_idx[i] 

569 _, curr_node_type, _ = self._batch_path[i][-1] 

570 relation, next_node_id = self._batch_curr_actions[i][act_idx] 

571 

572 if relation == "self_loop": 

573 next_node_type = curr_node_type 

574 else: 

575 next_node_type = self._get_next_node_type(curr_node_type, relation) 

576 self._batch_path[i].append((relation, next_node_type, next_node_id)) 

577 

578 self._done = self._is_done() # must run before get actions, etc. 

579 self._batch_curr_state = self._batch_get_state(self._batch_path) 

580 self._batch_curr_actions = self._batch_get_actions(self._batch_path, self._done) 

581 self._batch_curr_reward = self._batch_get_reward(self._batch_path) 

582 return self._batch_curr_state, self._batch_curr_reward, self._done 

583 

584 def batch_action_mask(self, dropout): 

585 # Return action masks of size [bs, act_dim]. 

586 batch_mask = [] 

587 for actions in self._batch_curr_actions: 

588 act_idxs = np.arange(len(actions)) 

589 if dropout > 0 and len(act_idxs) >= 5: # noqa: PLR2004 

590 keep_size = int(len(act_idxs[1:]) * (1.0 - dropout)) 

591 tmp = self.rng.choice(act_idxs[1:], keep_size, replace=False).tolist() 

592 act_idxs = np.concatenate([[act_idxs[0]], tmp]) 

593 act_mask = np.zeros(self.act_dim, dtype=np.uint8) 

594 act_mask[act_idxs] = 1 

595 batch_mask.append(act_mask) 

596 return np.vstack(batch_mask) 

597 

598 def _batch_acts_to_masks(self, batch_acts): 

599 batch_masks = np.zeros((len(batch_acts), self.act_dim), dtype=np.uint8) 

600 for i, acts in enumerate(batch_acts): 

601 num_acts = len(acts) 

602 batch_masks[i, :num_acts] = 1 

603 return batch_masks 

604 

605 def predict(self, interaction): 

606 if self.train_stage == "pretrain": 

607 return self.predict_transe(interaction) 

608 return 

609 

610 def full_sort_predict(self, interaction): 

611 if self.train_stage == "pretrain": 

612 return self.full_sort_predict_transe(interaction) 

613 

614 # get set temporal weights 

615 interaction, temporal_weight = interaction 

616 users = interaction[self.USER_ID] 

617 paths, probs = self.beam_search(users) 

618 interacted_matrix = self._build_interacted_matrix(temporal_weight) 

619 return self.collect_scores(users, paths, probs, interacted_matrix) 

620 

621 def _build_interacted_matrix(self, temporal_weight): 

622 item_emb = None 

623 purchase_matrix = [] 

624 for uid in range(1, len(temporal_weight.uc_weight) + 1): 

625 for clus_wt in temporal_weight.uc_weight[uid]: 

626 if item_emb is None: 

627 item_emb = self.ui_clust_relation_embedding[clus_wt] * temporal_weight.uc_weight[uid][clus_wt] 

628 else: 

629 item_emb += self.ui_clust_relation_embedding[clus_wt] * temporal_weight.uc_weight[uid][clus_wt] 

630 purchase_matrix.append(item_emb) 

631 

632 return purchase_matrix 

633 

634 def explain(self, interaction): 

635 """Support function used for case study. 

636 

637 Args: 

638 interaction : test interaction data 

639 

640 Returns: 

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

642 """ 

643 users, temporal_weight = interaction[self.USER_ID] 

644 paths, probs = self.beam_search(users) 

645 interacted_matrix = self._build_interacted_matrix(temporal_weight) 

646 

647 scores, explanations = self.collect_scores(users, paths, probs, interacted_matrix) 

648 

649 for exp in explanations: 

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

651 

652 return scores, explanations 

653 

654 def decode_path(self, path): 

655 return path 

656 

657 def beam_search(self, users): 

658 users = [user.item() for user in users] 

659 state_pool = self.reset(users) # numpy of [bs, dim] 

660 path_pool = self._batch_path # list of list, size=bs 

661 probs_pool = [[] for _ in users] 

662 

663 for hop, k in enumerate(self.beam_search_hop): 

664 state_tensor = torch.FloatTensor(state_pool).to(self.device) 

665 acts_pool = self._batch_get_actions(path_pool, False) # list of list, size=bs 

666 actmask_pool = self._batch_acts_to_masks(acts_pool) # numpy of [bs, dim] 

667 actmask_tensor = torch.BoolTensor(actmask_pool).to(self.device) 

668 # Tensor of [bs, act_dim] 

669 probs, _ = self.forward((state_tensor, actmask_tensor)) 

670 # In order to differ from masked actions 

671 probs = probs + actmask_tensor.float() 

672 topk_probs, topk_idxs = torch.topk(probs, k, dim=1) # LongTensor of [bs, k] 

673 

674 topk_idxs = topk_idxs.detach().cpu().numpy() 

675 topk_probs = topk_probs.detach().cpu().numpy() 

676 

677 new_path_pool, new_probs_pool = [], [] 

678 for row in range(topk_idxs.shape[0]): 

679 path = path_pool[row] 

680 probs = probs_pool[row] 

681 for idx, p in zip(topk_idxs[row], topk_probs[row]): 

682 if idx >= len(acts_pool[row]): # act idx is invalid 

683 continue 

684 

685 # (relation, next_node_id) 

686 relation, next_node_id = acts_pool[row][idx] 

687 

688 if relation == "self_loop": 

689 next_node_type = path[-1][1] 

690 else: 

691 next_node_type = self._get_next_node_type(path[-1][1], relation) 

692 

693 new_path = path + [(relation, next_node_type, next_node_id)] 

694 

695 new_path_pool.append(new_path) 

696 new_probs_pool.append(probs + [p]) 

697 path_pool = new_path_pool 

698 probs_pool = new_probs_pool 

699 if hop < 2: # noqa: PLR2004 

700 state_pool = self._batch_get_state(path_pool) 

701 

702 return path_pool, probs_pool 

703 

704 def collect_scores(self, users, paths, probs, interacted_matrix): 

705 collect_results = list() 

706 pad_emb = np.zeros((1, self.embedding_size)) 

707 interacted_embeds = np.concatenate((pad_emb, np.array(interacted_matrix))) 

708 # 1) get all valid paths for each user, compute path score and path probability 

709 pred_paths = {uid.item(): defaultdict(list) for uid in users} 

710 path_scores = np.dot(self.user_embedding + interacted_embeds, self.entity_embedding[: self.n_items].T) 

711 for path, prob in zip(paths, probs): 

712 if "self_loop" in [node[0] for node in path[1:]]: 

713 continue 

714 

715 if path[-1][1] != "entity": 

716 continue 

717 

718 path_uid = path[0][2] 

719 # check it is a user in the test set batch 

720 if path_uid not in pred_paths: 

721 continue 

722 

723 path_pid = path[-1][2] 

724 # check it is an item 

725 if not (path_pid < self.n_items): 

726 continue 

727 

728 if path_pid in self.positives[path_uid]: 

729 continue 

730 

731 path_score = path_scores[path_uid][path_pid] 

732 path_prob = reduce(lambda x, y: x * y, prob) 

733 pred_paths[path_uid][path_pid].append((path_score, path_prob, path)) 

734 

735 # 2) Pick best paths for each user-item pair based on the score 

736 best_pred_paths = defaultdict(list) 

737 for user, user_pred_paths in pred_paths.items(): 

738 for item in user_pred_paths: 

739 if item in self.positives[user]: 

740 continue 

741 # Get the path with highest probability 

742 sorted_path = sorted(user_pred_paths[item], key=lambda x: x[1], reverse=True)[0] 

743 best_pred_paths[user].append(sorted_path) 

744 

745 # 3) Fill the results tensor 

746 results = torch.full((len(users), self.n_items), -torch.inf) 

747 

748 for i, user in enumerate(best_pred_paths): 

749 # sort by score 

750 sorted_path = sorted(best_pred_paths[user], key=lambda x: (x[0], x[1]), reverse=True) 

751 top_items = [[p[-1][2], score] for score, _, p in sorted_path][: max(self.topk)] 

752 top_paths = [p for _, _, p in sorted_path][: max(self.topk)] 

753 if len(top_items) < max(self.topk): 

754 cand_pids = np.argsort(path_scores[user]) 

755 for cand_pids in cand_pids[::-1]: 

756 if cand_pids in self.positives[user]: 

757 continue 

758 top_items.append([cand_pids, path_scores[user][cand_pids]]) 

759 if len(top_items) >= max(self.topk): 

760 break 

761 # Change order from smallest to largest 

762 top_items = top_items[::-1] 

763 top_paths = top_paths[::-1] 

764 for (item, score), path in zip(top_items, top_paths): 

765 results[i, item] = score.tolist() 

766 

767 # collect user, item, score and paths. 

768 collect_results.append([user, item, score, path]) 

769 

770 return results, collect_results 

771 

772 

773class KGState: 

774 def __init__(self, embedding_size, history_len): 

775 self.embedding_size = embedding_size 

776 self.history_len = history_len # mode: one of {full, current} 

777 if history_len == 0: 

778 self.dim = 2 * embedding_size 

779 elif history_len == 1: 

780 self.dim = 4 * embedding_size 

781 elif history_len == 2: # noqa: PLR2004 

782 self.dim = 6 * embedding_size 

783 else: 

784 raise Exception("history length should be one of {0, 1, 2}") 

785 

786 def __call__( 

787 self, user_embed, node_embed, last_node_embed, last_relation_embed, older_node_embed, older_relation_embed 

788 ): 

789 if self.history_len == 0: 

790 return np.concatenate([user_embed, node_embed]) 

791 elif self.history_len == 1: 

792 return np.concatenate([user_embed, node_embed, last_node_embed, last_relation_embed]) 

793 elif self.history_len == 2: # noqa: PLR2004 

794 return np.concatenate( 

795 [user_embed, node_embed, last_node_embed, last_relation_embed, older_node_embed, older_relation_embed] 

796 ) 

797 else: 

798 raise ValueError("mode should be one of {full, current}")