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
« 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"""PGPR
6##################################################
7Reference:
8 Xian et al. "Reinforcement Knowledge Graph Reasoning for Explainable Recommendation." in SIGIR 2019.
10Reference code:
11 https://github.com/orcax/PGPR
12"""
14from collections import defaultdict, namedtuple
15from functools import reduce
17import numpy as np
18import torch
19import torch.nn.functional as F
20from torch import nn
21from torch.distributions import Categorical
23from hopwise.model.abstract_recommender import ExplainableRecommender, KnowledgeRecommender
24from hopwise.utils import InputType
27class PGPR(KnowledgeRecommender, ExplainableRecommender):
28 input_type = InputType.USERWISE
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"]
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"]
49 self.fix_scores_sorting_bug = config["fix_scores_sorting_bug"]
51 # user-item relation
52 self.ui_relation_id = dataset.field2token_id["relation_id"][dataset.ui_relation]
54 self.graph_dict = dataset.ckg_dict_graph()
56 # Items
57 self.positives = dataset.history_item_matrix()[0]
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)
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)
72 # Self Loop Embedding
73 self.self_loop_embedding = np.zeros(self.embedding_size)
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()}
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 }
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)
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 = []
115 # random generator
116 self.rng = np.random.default_rng()
118 self.SavedAction = namedtuple("SavedAction", ["log_prob", "value"])
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
133 self.saved_actions.append(self.SavedAction(m.log_prob(acts), value))
134 self.entropy.append(m.entropy())
135 return acts.cpu().numpy().tolist()
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
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]
149 for i in range(1, num_steps):
150 batch_rewards[:, num_steps - i - 1] += self.gamma * batch_rewards[:, num_steps - i]
152 actor_loss = 0
153 critic_loss = 0
154 entropy_loss = 0
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, ]
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
169 del self.rewards[:]
170 del self.saved_actions[:]
171 del self.entropy[:]
173 return loss, actor_loss, critic_loss, entropy_loss
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
189 def calculate_loss(self, interaction):
190 users = interaction[self.USER_ID]
191 users = users[users != 0]
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()
203 return loss
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
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"
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
223 def _batch_get_actions(self, batch_path, done):
224 return [self._get_actions(path, done) for path in batch_path]
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)]
231 # (1) If game is finished, only return self-loop action.
232 if done:
233 return actions
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()
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])
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))
253 # (3) If candidate action set is empty, only return self-loop action.
254 if len(candidate_acts) == 0:
255 return actions
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
263 # (5) If there are too many actions, do some deterministic trimming here!
264 user_embed = self.user_embedding[path[0][-1]]
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]
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)
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
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]
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)
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
305 older_relation, last_node_type, last_node_id = path[-2]
306 last_relation, curr_node_type, curr_node_id = path[-1]
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]
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
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
322 _, older_node_type, older_node_id = path[-3]
323 older_node_embed = self.node_type2emb[older_node_type][older_node_id]
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]
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
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)
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
344 if not self._has_pattern(path):
345 return 0.0
347 target_score = 0.0
348 _, curr_node_type, curr_node_id = path[-1]
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
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
363 def batch_step(self, batch_act_idx):
364 assert len(batch_act_idx) == len(self._batch_path)
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]
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))
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
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)
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
405 def predict(self, interaction):
406 return
408 def full_sort_predict(self, interaction):
409 users = interaction[self.USER_ID]
411 paths, probs = self.beam_search(users)
413 scores, _ = self.collect_scores(users, paths, probs)
415 return scores
417 def explain(self, interaction):
418 """Support function used for case study.
420 Args:
421 interaction : test interaction data
423 Returns:
424 pd.Dataframe: explanation results with columns: "user", "item", "score", "path"
425 """
426 users = interaction[self.USER_ID]
428 paths, probs = self.beam_search(users)
430 scores, explanations = self.collect_scores(users, paths, probs)
432 for exp in explanations:
433 exp[-1] = self.decode_path(exp[-1])
435 return scores, explanations
437 def decode_path(self, path):
438 return path
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]
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]
456 topk_idxs = topk_idxs.detach().cpu().numpy()
457 topk_probs = topk_probs.detach().cpu().numpy()
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
467 # (relation, next_node_id)
468 relation, next_node_id = acts_pool[row][idx]
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)
475 new_path = path + [(relation, next_node_type, next_node_id)]
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)
484 return path_pool, probs_pool
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
497 if path[-1][1] != "entity":
498 continue
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
505 path_pid = path[-1][2]
506 # check it is an item
507 if not (path_pid < self.n_items):
508 continue
510 if path_pid in self.positives[path_uid]:
511 continue
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))
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)
527 # 3) Fill the results tensor
528 results = torch.full((len(users), self.n_items), -torch.inf)
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
549 # collect user, item, score and paths.
550 collect_results.append([user, item, score, path])
552 return results, collect_results
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}")
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}")