Coverage for hopwise/model/knowledge_aware_recommender/pgpr.py: 89%

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

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

7Reference: 

8 Xian et al. "Reinforcement Knowledge Graph Reasoning for Explainable Recommendation." in SIGIR 2019. 

9 

10Reference code: 

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

12""" 

13 

14from collections import defaultdict, namedtuple 

15from functools import reduce 

16 

17import numpy as np 

18import torch 

19import torch.nn.functional as F 

20from torch import nn 

21from torch.distributions import Categorical 

22 

23from hopwise.model.abstract_recommender import ExplainableRecommender, KnowledgeRecommender 

24from hopwise.utils import InputType 

25 

26 

27class PGPR(KnowledgeRecommender, ExplainableRecommender): 

28 input_type = InputType.USERWISE 

29 

30 def __init__(self, config, dataset): 

31 super().__init__(config, dataset) 

32 # Load parameters info from config 

33 self.user_num = dataset.user_num 

34 self.device = config["device"] 

35 self.topk = config["topk"] 

36 

37 # PGPR Configurations 

38 self.state_history = config["state_history"] 

39 self.max_acts = config["max_acts"] 

40 self.gamma = config["gamma"] 

41 self.action_dropout = config["action_dropout"] 

42 self.hidden_sizes = config["hidden_sizes"] 

43 self.act_dim = self.max_acts + 1 

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

45 self.weight_factor = config["weight_factor"] 

46 self.path_pattern = config["path_constraint"] 

47 self.beam_search_hop = config["beam_search_hop"] 

48 

49 self.fix_scores_sorting_bug = config["fix_scores_sorting_bug"] 

50 

51 # user-item relation 

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

53 

54 self.graph_dict = dataset.ckg_dict_graph() 

55 

56 # Items 

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

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 self.embedding_size = self.user_embedding.shape[1] 

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

65 

66 # Actor-Critic model 

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

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

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

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

71 

72 # Self Loop Embedding 

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

74 

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

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

77 

78 # Mapping node type to embeddings 

79 self.node_type2emb = { 

80 "user": self.user_embedding, 

81 "entity": self.entity_embedding, 

82 "relation": self.relation_embedding, 

83 "self_loop": self.self_loop_embedding, 

84 } 

85 

86 # Normalization score 

87 u_p_scores = np.dot( 

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

89 ) 

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

91 

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

93 # Preprocess the path constraints. 

94 self.patterns = list() 

95 for path_constraint in self.path_pattern: 

96 relations = list() 

97 for node in path_constraint: 

98 path_rel = node[0] 

99 if path_rel is not None: 

100 if path_rel.endswith("_r"): 

101 # remove the reverse suffix 

102 path_rel = path_rel[:-2] 

103 relations.append(path_rel) 

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

105 # Following is current episode information. 

106 self._batch_path = None 

107 self._batch_curr_actions = None 

108 self._batch_curr_state = None 

109 self._batch_curr_reward = None 

110 self._done = False 

111 self.saved_actions = [] 

112 self.rewards = [] 

113 self.entropy = [] 

114 

115 # random generator 

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

117 

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

119 

120 def select_action(self, batch_state, batch_act_mask): 

121 # Tensor [bs, state_dim] 

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

123 # Tensor of [bs, act_dim] 

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

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

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

127 m = Categorical(probs) 

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

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

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

131 acts[valid_idx == 0] = 0 

132 

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

134 self.entropy.append(m.entropy()) 

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

136 

137 def update(self): # prev update 

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

139 del self.rewards[:] 

140 del self.saved_actions[:] 

141 del self.entropy[:] 

142 return 0.0, 0.0, 0.0 

143 

144 # numpy array of [bs, #steps] 

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

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

147 num_steps = batch_rewards.shape[1] 

148 

149 for i in range(1, num_steps): 

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

151 

152 actor_loss = 0 

153 critic_loss = 0 

154 entropy_loss = 0 

155 

156 for i in range(0, num_steps): 

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

158 log_prob, value = self.saved_actions[i] 

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

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

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

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

163 

164 actor_loss = actor_loss.mean() 

165 critic_loss = critic_loss.mean() 

166 entropy_loss = entropy_loss.mean() 

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

168 

169 del self.rewards[:] 

170 del self.saved_actions[:] 

171 del self.entropy[:] 

172 

173 return loss, actor_loss, critic_loss, entropy_loss 

174 

175 def forward(self, inputs): 

176 # used only in inference 

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

178 state, act_mask = inputs 

179 x = self.l1(state) 

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

181 out = self.l2(x) 

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

183 actor_logits = self.actor(x) 

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

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

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

187 return act_probs, state_values 

188 

189 def calculate_loss(self, interaction): 

190 users = interaction[self.USER_ID] 

191 users = users[users != 0] 

192 

193 # Policy training 

194 batch_state = self.reset(users) 

195 done = False 

196 while not done: 

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

198 batch_act_idx = self.select_action(batch_state, batch_act_mask) 

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

200 self.rewards.append(batch_reward) 

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

202 

203 return loss 

204 

205 def _has_pattern(self, path): 

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

207 return pattern in self.patterns 

208 

209 def _get_next_node_type(self, current_node_type, relation_id): 

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

211 return "user" 

212 else: 

213 return "entity" 

214 

215 def reset(self, user): 

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

217 self._done = False 

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

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

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

221 return self._batch_curr_state 

222 

223 def _batch_get_actions(self, batch_path, done): 

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

225 

226 def _get_actions(self, path, done): 

227 # Compute actions for current node. 

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

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

230 

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

232 if done: 

233 return actions 

234 

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

236 # [CAVEAT] Must remove visited nodes! 

237 if isinstance(curr_node_id, torch.Tensor): 

238 curr_node_id = curr_node_id.item() 

239 

240 try: 

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

242 except KeyError: 

243 relations_nodes = [] 

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

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

246 

247 for r in relations_nodes: 

248 next_node_type = self._get_next_node_type(curr_node_type, r) 

249 next_node_ids = relations_nodes[r] 

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

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

252 

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

254 if len(candidate_acts) == 0: 

255 return actions 

256 

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

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

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

260 actions.extend(candidate_acts) 

261 return actions 

262 

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

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

265 

266 scores = [] 

267 for r, next_node_id in candidate_acts: 

268 next_node_type = self._get_next_node_type(curr_node_type, r) 

269 if next_node_type == "user": 

270 src_embed = user_embed 

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

272 src_embed = user_embed + self.relation_embedding[self.ui_relation_id] 

273 else: 

274 src_embed = user_embed + self.relation_embedding[self.ui_relation_id] + self.relation_embedding[r] 

275 

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

277 # This trimming may filter out target items! 

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

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

280 # score = 99999.0 

281 scores.append(score) 

282 

283 # choose actions with larger scores 

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

285 if self.fix_scores_sorting_bug: 

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

287 else: 

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

289 actions.extend(candidate_acts) 

290 return actions 

291 

292 def _batch_get_state(self, batch_path): 

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

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

295 

296 def _get_state(self, path): 

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

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

299 zero_embed = np.zeros(self.embedding_size) 

300 

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

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

303 return state 

304 

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

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

307 

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

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

310 

311 if last_relation != "self_loop": 

312 last_relation_embed = self.relation_embedding[last_relation] 

313 else: 

314 last_relation_embed = self.self_loop_embedding 

315 

316 if len(path) == 2: 

317 state = self.state_gen( 

318 user_embed, curr_node_embed, last_node_embed, last_relation_embed, zero_embed, zero_embed 

319 ) 

320 return state 

321 

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

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

324 

325 if older_relation == "self_loop": 

326 older_relation_embed = self.self_loop_embedding 

327 else: 

328 older_relation_embed = self.relation_embedding[older_relation] 

329 

330 state = self.state_gen( 

331 user_embed, curr_node_embed, last_node_embed, last_relation_embed, older_node_embed, older_relation_embed 

332 ) 

333 return state 

334 

335 def _batch_get_reward(self, batch_path): 

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

337 return np.array(batch_reward) 

338 

339 def _get_reward(self, path): 

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

341 if len(path) <= 2: 

342 return 0.0 

343 

344 if not self._has_pattern(path): 

345 return 0.0 

346 

347 target_score = 0.0 

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

349 

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

351 # Give soft reward for other reached items. 

352 uid = path[0][-1] 

353 u_vec = self.user_embedding[uid] + self.relation_embedding[self.ui_relation_id] 

354 p_vec = self.entity_embedding[curr_node_id] 

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

356 target_score = max(score, 0.0) 

357 return target_score 

358 

359 def _is_done(self): 

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

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

362 

363 def batch_step(self, batch_act_idx): 

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

365 

366 # Execute batch actions. 

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

368 act_idx = batch_act_idx[i] 

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

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

371 

372 if relation == "self_loop": 

373 next_node_type = curr_node_type 

374 else: 

375 next_node_type = self._get_next_node_type(curr_node_type, relation) 

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

377 

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

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

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

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

382 return self._batch_curr_state, self._batch_curr_reward, self._done 

383 

384 def batch_action_mask(self, dropout): 

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

386 batch_mask = [] 

387 for actions in self._batch_curr_actions: 

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

389 if dropout > 0 and len(act_idxs) >= 5: 

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

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

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

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

394 act_mask[act_idxs] = 1 

395 batch_mask.append(act_mask) 

396 return np.vstack(batch_mask) 

397 

398 def _batch_acts_to_masks(self, batch_acts): 

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

400 for i, acts in enumerate(batch_acts): 

401 num_acts = len(acts) 

402 batch_masks[i, :num_acts] = 1 

403 return batch_masks 

404 

405 def predict(self, interaction): 

406 return 

407 

408 def full_sort_predict(self, interaction): 

409 users = interaction[self.USER_ID] 

410 

411 paths, probs = self.beam_search(users) 

412 

413 scores, _ = self.collect_scores(users, paths, probs) 

414 

415 return scores 

416 

417 def explain(self, interaction): 

418 """Support function used for case study. 

419 

420 Args: 

421 interaction : test interaction data 

422 

423 Returns: 

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

425 """ 

426 users = interaction[self.USER_ID] 

427 

428 paths, probs = self.beam_search(users) 

429 

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

431 

432 for exp in explanations: 

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

434 

435 return scores, explanations 

436 

437 def decode_path(self, path): 

438 return path 

439 

440 def beam_search(self, users): 

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

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

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

444 probs_pool = [[] for _ in users] 

445 

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

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

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

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

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

451 probs, _ = self.forward((state_tensor, actmask_tensor)) # Tensor of [bs, act_dim] 

452 # In order to differ from masked actions 

453 probs = probs + actmask_tensor.float() 

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

455 

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

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

458 

459 new_path_pool, new_probs_pool = [], [] 

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

461 path = path_pool[row] 

462 probs = probs_pool[row] 

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

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

465 continue 

466 

467 # (relation, next_node_id) 

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

469 

470 if relation == "self_loop": 

471 next_node_type = path[-1][1] 

472 else: 

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

474 

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

476 

477 new_path_pool.append(new_path) 

478 new_probs_pool.append(probs + [p]) 

479 path_pool = new_path_pool 

480 probs_pool = new_probs_pool 

481 if hop < 2: 

482 state_pool = self._batch_get_state(path_pool) 

483 

484 return path_pool, probs_pool 

485 

486 def collect_scores(self, users, paths, probs): 

487 collect_results = list() 

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

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

490 path_scores = np.dot( 

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

492 ) 

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

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

495 continue 

496 

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

498 continue 

499 

500 path_uid = path[0][2] 

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

502 if path_uid not in pred_paths: 

503 continue 

504 

505 path_pid = path[-1][2] 

506 # check it is an item 

507 if not (path_pid < self.n_items): 

508 continue 

509 

510 if path_pid in self.positives[path_uid]: 

511 continue 

512 

513 path_score = path_scores[path_uid][path_pid] 

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

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

516 

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

518 best_pred_paths = defaultdict(list) 

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

520 for item in user_pred_paths: 

521 if item in self.positives[user]: 

522 continue 

523 # Get the path with highest probability 

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

525 best_pred_paths[user].append(sorted_path) 

526 

527 # 3) Fill the results tensor 

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

529 

530 for i, user in enumerate(best_pred_paths): 

531 # sort by score 

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

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

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

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

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

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

538 if cand_pids in self.positives[user]: 

539 continue 

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

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

542 break 

543 # Change order from smallest to largest 

544 top_items = top_items[::-1] 

545 top_paths = top_paths[::-1] 

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

547 results[i, item] = score 

548 

549 # collect user, item, score and paths. 

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

551 

552 return results, collect_results 

553 

554 

555class KGState: 

556 def __init__(self, embedding_size, history_len): 

557 self.embedding_size = embedding_size 

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

559 if history_len == 0: 

560 self.dim = 2 * embedding_size 

561 elif history_len == 1: 

562 self.dim = 4 * embedding_size 

563 elif history_len == 2: 

564 self.dim = 6 * embedding_size 

565 else: 

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

567 

568 def __call__( 

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

570 ): 

571 if self.history_len == 0: 

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

573 elif self.history_len == 1: 

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

575 elif self.history_len == 2: 

576 return np.concatenate( 

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

578 ) 

579 else: 

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