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
« 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
5"""TPRec
6##################################################
7Reference: Time-aware Path Reasoning on Knowledge Graph for Recommendation (https://arxiv.org/pdf/2108.02634)
9"""
11from collections import defaultdict, namedtuple
12from functools import reduce
14import numpy as np
15import torch
16import torch.nn.functional as F
17from torch import nn
18from torch.distributions import Categorical
20from hopwise.model.abstract_recommender import ExplainableRecommender, KnowledgeRecommender
21from hopwise.utils import InputType
24class TPRec(KnowledgeRecommender, ExplainableRecommender):
25 """
26 TPRec
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
32 """
34 input_type = InputType.USERWISE
36 def __init__(self, config, dataset):
37 super().__init__(config, dataset)
38 self.config = config
40 # Load parameters info from config
41 self.user_num = dataset.user_num
42 self.device = config["device"]
43 self.topk = config["topk"]
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"]
61 # user-item relation
62 self.ui_relation = dataset.ui_relation
63 self.ui_relation_id = dataset.field2token_id["relation_id"][dataset.ui_relation]
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
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")
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)
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)
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)
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)
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)
113 # Self Loop Embedding
114 self.self_loop_embedding = np.zeros(self.embedding_size)
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()}
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 }
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)
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))
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
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 = []
164 # random generator
165 self.rng = np.random.default_rng()
166 self.SavedAction = namedtuple("SavedAction", ["log_prob", "value"])
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.
173 Args:
174 path_constraint in the form of [rel1, rel2, rel3]
175 """
177 def _expand_recursive(current_path, remaining_relations):
178 # Base case - path complete
179 if not remaining_relations:
180 return [current_path]
182 # Get next relation
183 relation = remaining_relations[0]
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:])
196 # Start expansion from empty path
197 return _expand_recursive([], path_constraint[1:]) # Skip self-loop
199 def _get_pretrained_weights(self):
200 import os
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
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
223 self.saved_actions.append(self.SavedAction(m.log_prob(acts), value))
224 self.entropy.append(m.entropy())
225 return acts.cpu().numpy().tolist()
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
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]
239 for i in range(1, num_steps):
240 batch_rewards[:, num_steps - i - 1] += self.gamma * batch_rewards[:, num_steps - i]
242 actor_loss = 0
243 critic_loss = 0
244 entropy_loss = 0
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, ]
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
259 del self.rewards[:]
260 del self.saved_actions[:]
261 del self.entropy[:]
263 return loss, actor_loss, critic_loss, entropy_loss
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
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)
286 return user_e, pos_item_e, neg_item_e, rec_r_e
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
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]
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)
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])
332 loss += self.transe_loss(h_e + r_e, pos_t_e, neg_t_e)
334 return loss
336 def forward_transe(self, user, relation, item):
337 score = -torch.norm(user + relation - item, p=2, dim=1)
338 return score
340 def predict_transe(self, interaction):
341 user = interaction[self.USER_ID]
342 item = interaction[self.ITEM_ID]
344 user_e = self.user_embedding(user)
345 item_e = self.entity_embedding(item)
347 rec_r_e = self.relation_embedding.weight[-1]
348 rec_r_e = rec_r_e.expand_as(user_e)
350 return self.forward_transe(user_e, rec_r_e, item_e)
352 def full_sort_predict_transe(self, interaction):
353 user = interaction[self.USER_ID]
354 user_e = self.user_embedding(user)
356 rec_r_e = self.relation_embedding.weight[-1]
357 rec_r_e = rec_r_e.expand_as(user_e)
359 item_indices = torch.tensor(range(self.n_items)).to(self.device)
360 all_item_e = self.entity_embedding.weight[item_indices]
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)
366 return -torch.norm(user_e + rec_r_e - t, p=2, dim=2)
368 def calculate_loss(self, interaction):
369 if self.train_stage == "pretrain":
370 return self.calculate_loss_transe(interaction)
372 users = interaction[self.USER_ID]
373 users = users[users != 0]
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()
385 return loss
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
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"
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
405 def _batch_get_actions(self, batch_path, done):
406 return [self._get_actions(path, done) for path in batch_path]
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)]
413 # (1) If game is finished, only return self-loop action.
414 if done:
415 return actions
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()
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])
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))
435 # (3) If candidate action set is empty, only return self-loop action.
436 if len(candidate_acts) == 0:
437 return actions
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
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]
449 scores = []
450 item_emb = None
451 if isinstance(uid, torch.Tensor):
452 uid = uid.item()
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]
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]
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)
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
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]
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)
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
498 older_relation, last_node_type, last_node_id = path[-2]
499 last_relation, curr_node_type, curr_node_id = path[-1]
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]
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
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
515 _, older_node_type, older_node_id = path[-3]
516 older_node_embed = self.node_type2emb[older_node_type][older_node_id]
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]
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
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)
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
537 if not self._has_pattern(path):
538 return 0.0
540 target_score = 0.0
541 _, curr_node_type, curr_node_id = path[-1]
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]
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
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
563 def batch_step(self, batch_act_idx):
564 assert len(batch_act_idx) == len(self._batch_path)
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]
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))
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
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)
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
605 def predict(self, interaction):
606 if self.train_stage == "pretrain":
607 return self.predict_transe(interaction)
608 return
610 def full_sort_predict(self, interaction):
611 if self.train_stage == "pretrain":
612 return self.full_sort_predict_transe(interaction)
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)
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)
632 return purchase_matrix
634 def explain(self, interaction):
635 """Support function used for case study.
637 Args:
638 interaction : test interaction data
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)
647 scores, explanations = self.collect_scores(users, paths, probs, interacted_matrix)
649 for exp in explanations:
650 exp[-1] = self.decode_path(exp[-1])
652 return scores, explanations
654 def decode_path(self, path):
655 return path
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]
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]
674 topk_idxs = topk_idxs.detach().cpu().numpy()
675 topk_probs = topk_probs.detach().cpu().numpy()
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
685 # (relation, next_node_id)
686 relation, next_node_id = acts_pool[row][idx]
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)
693 new_path = path + [(relation, next_node_type, next_node_id)]
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)
702 return path_pool, probs_pool
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
715 if path[-1][1] != "entity":
716 continue
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
723 path_pid = path[-1][2]
724 # check it is an item
725 if not (path_pid < self.n_items):
726 continue
728 if path_pid in self.positives[path_uid]:
729 continue
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))
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)
745 # 3) Fill the results tensor
746 results = torch.full((len(users), self.n_items), -torch.inf)
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()
767 # collect user, item, score and paths.
768 collect_results.append([user, item, score, path])
770 return results, collect_results
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}")
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}")