Coverage for hopwise/evaluator/collector.py: 59%
162 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 : 2021/6/23
2# @Author : Zihan Lin
3# @Email : zhlin@ruc.edu.cn
5# UPDATE
6# @Time : 2021/7/18
7# @Author : Zhichao Feng
8# @email : fzcbupt@gmail.com
10"""hopwise.evaluator.collector
11################################################
12"""
14import copy
16import torch
18from hopwise.evaluator.register import Register, Register_KG
19from hopwise.evaluator.utils import train_tsne
22class DataStruct:
23 def __init__(self):
24 self._data_dict = {}
26 def __getitem__(self, name: str):
27 return self._data_dict[name]
29 def __setitem__(self, name: str, value):
30 self._data_dict[name] = value
32 def __delitem__(self, name: str):
33 self._data_dict.pop(name)
35 def __contains__(self, key: str):
36 return key in self._data_dict
38 def get(self, name: str):
39 if name not in self._data_dict:
40 raise IndexError("Can not load the data without registration !")
41 return self[name]
43 def set(self, name: str, value):
44 self._data_dict[name] = value
46 def update_tensor(self, name: str, value):
47 if name not in self._data_dict:
48 if isinstance(value, torch.Tensor):
49 self._data_dict[name] = value.clone().detach()
50 else:
51 self._data_dict[name] = value
52 elif isinstance(self._data_dict[name], torch.Tensor):
53 self._data_dict[name] = torch.cat((self._data_dict[name], value.clone().detach()), dim=0)
54 else:
55 self._data_dict[name] = self._data_dict[name] + value
57 def __str__(self):
58 data_info = "\nContaining:\n"
59 for data_key in self._data_dict.keys():
60 data_info += data_key + "\n"
61 return data_info
64class Collector:
65 """The collector is used to collect the resource for evaluator.
66 As the evaluation metrics are various, the needed resource not only contain the recommended result
67 but also other resource from data and model. They all can be collected by the collector during the training
68 and evaluation process.
70 This class is only used in Trainer.
72 """
74 def __init__(self, config):
75 self.config = config
76 self.data_struct = DataStruct()
77 self.register = Register(config)
78 self.full = "full" in config["eval_args"]["mode"]
79 self.topk = self.config["topk"]
80 self.device = self.config["device"]
82 def train_data_collect(self, train_data):
83 """Collect the evaluation resource from training data.
85 Args:
86 train_data (AbstractDataLoader): the training dataloader which contains the training data.
88 """
89 if self.register.need("data.num_items"):
90 item_id = self.config["ITEM_ID_FIELD"]
91 self.data_struct.set("data.num_items", train_data.dataset.num(item_id))
92 if self.register.need("data.num_users"):
93 user_id = self.config["USER_ID_FIELD"]
94 self.data_struct.set("data.num_users", train_data.dataset.num(user_id))
95 if self.register.need("data.count_items"):
96 self.data_struct.set("data.count_items", train_data.dataset.item_counter)
97 if self.register.need("data.count_users"):
98 self.data_struct.set("data.count_users", train_data.dataset.user_counter)
99 if self.register.need("data.history_index"):
100 row = train_data.dataset.inter_feat[train_data.dataset.uid_field]
101 col = train_data.dataset.inter_feat[train_data.dataset.iid_field]
102 self.data_struct.set("data.history_index", torch.vstack([row, col]))
103 if self.register.need("data.timestamp"):
104 temporal_matrix = train_data.dataset.inter_matrix(value_field=train_data.dataset.time_field).toarray()
105 self.data_struct.set("data.timestamp", temporal_matrix)
107 def eval_data_collect(self, eval_data):
108 """Collect the evaluation resource from evaluation data, such as user and item features.
110 Args:
111 eval_data (AbstractDataLoader): the evaluation dataloader which contains the evaluation data.
113 """
114 if self.register.need("eval_data.user_feat"):
115 if not hasattr(eval_data.dataset, "user_feat") or eval_data.dataset.user_feat is None:
116 raise AttributeError("Evaluation data does not include user features.")
117 self.data_struct.set("eval_data.user_feat", eval_data.dataset.user_feat)
119 def _average_rank(self, scores):
120 """Get the ranking of an ordered tensor, and take the average of the ranking for positions with equal values.
122 Args:
123 scores(tensor): an ordered tensor, with size of `(N, )`
125 Returns:
126 torch.Tensor: average_rank
128 Example:
129 >>> average_rank(tensor([[1,2,2,2,3,3,6],[2,2,2,2,4,5,5]]))
130 tensor([[1.0000, 3.0000, 3.0000, 3.0000, 5.5000, 5.5000, 7.0000],
131 [2.5000, 2.5000, 2.5000, 2.5000, 5.0000, 6.5000, 6.5000]])
133 Reference:
134 https://github.com/scipy/scipy/blob/v0.17.1/scipy/stats/stats.py#L5262-L5352
136 """
137 length, width = scores.shape
138 true_tensor = torch.full((length, 1), True, dtype=torch.bool, device=self.device)
140 obs = torch.cat([true_tensor, scores[:, 1:] != scores[:, :-1]], dim=1)
141 # bias added to dense
142 bias = torch.arange(0, length, device=self.device).repeat(width).reshape(width, -1).transpose(1, 0).reshape(-1)
143 dense = obs.view(-1).cumsum(0) + bias
145 # cumulative counts of each unique value
146 count = torch.where(torch.cat([obs, true_tensor], dim=1))[1]
147 # get average rank
148 avg_rank = 0.5 * (count[dense] + count[dense - 1] + 1).view(length, -1)
150 return avg_rank
152 def eval_batch_collect(
153 self,
154 scores,
155 interaction,
156 positive_u: torch.Tensor,
157 positive_i: torch.Tensor,
158 ):
159 """Collect the evaluation resource from batched eval data and batched model output.
161 Args:
162 scores (Torch.Tensor): the output tensor of model with the shape of `(N, )`
163 interaction (Interaction): batched eval data.
164 positive_u (Torch.Tensor): the row index of positive items for each user.
165 positive_i (Torch.Tensor): the positive item id for each user.
166 """
167 if self.register.need("rec.users"):
168 uid_field = self.config["USER_ID_FIELD"]
169 self.data_struct.update_tensor("rec.users", interaction[uid_field])
171 if self.register.need("rec.items"):
172 # get topk
173 _, topk_idx = torch.topk(scores, max(self.topk), dim=-1) # n_users x k
174 self.data_struct.update_tensor("rec.items", topk_idx)
176 if self.register.need("rec.topk"):
177 _, topk_idx = torch.topk(scores, max(self.topk), dim=-1) # n_users x k
178 pos_matrix = torch.zeros_like(scores, dtype=torch.int)
179 pos_matrix[positive_u, positive_i] = 1
180 pos_len_list = pos_matrix.sum(dim=1, keepdim=True)
181 pos_idx = torch.gather(pos_matrix, dim=1, index=topk_idx)
182 result = torch.cat((pos_idx, pos_len_list), dim=1)
183 self.data_struct.update_tensor("rec.topk", result)
185 if self.register.need("rec.meanrank"):
186 desc_scores, desc_index = torch.sort(scores, dim=-1, descending=True)
188 # get the index of positive items in the ranking list
189 pos_matrix = torch.zeros_like(scores)
190 pos_matrix[positive_u, positive_i] = 1
191 pos_index = torch.gather(pos_matrix, dim=1, index=desc_index)
193 avg_rank = self._average_rank(desc_scores)
194 pos_rank_sum = torch.where(pos_index == 1, avg_rank, torch.zeros_like(avg_rank)).sum(dim=-1, keepdim=True)
196 pos_len_list = pos_matrix.sum(dim=1, keepdim=True)
197 user_len_list = desc_scores.argmin(dim=1, keepdim=True)
198 result = torch.cat((pos_rank_sum, user_len_list, pos_len_list), dim=1)
199 self.data_struct.update_tensor("rec.meanrank", result)
201 if self.register.need("rec.score"):
202 self.data_struct.update_tensor("rec.score", scores)
204 if self.register.need("data.label"):
205 self.label_field = self.config["LABEL_FIELD"]
206 self.data_struct.update_tensor("data.label", interaction[self.label_field].to(self.device))
208 def model_collect(self, model: torch.nn.Module, load_best_model=False):
209 """Collect the evaluation resource from model and do something with the model.
211 Args:
212 model (nn.Module): the trained recommendation model.
213 load_best_model (bool): whether to load the best model.
214 """
216 if self.config["tsne"] is not None:
217 train_tsne(model, self.config["tsne"], load_best_model)
219 def eval_collect(self, eval_pred: torch.Tensor, data_label: torch.Tensor):
220 """Collect the evaluation resource from total output and label.
221 It was designed for those models that can not predict with batch.
223 Args:
224 eval_pred (torch.Tensor): the output score tensor of model.
225 data_label (torch.Tensor): the label tensor.
226 """
227 if self.register.need("rec.score"):
228 self.data_struct.update_tensor("rec.score", eval_pred)
230 if self.register.need("data.label"):
231 self.label_field = self.config["LABEL_FIELD"]
232 self.data_struct.update_tensor("data.label", data_label.to(self.device))
234 def get_data_struct(self):
235 """Get all the evaluation resource that been collected.
236 And reset some of outdated resource.
237 """
238 for key in self.data_struct._data_dict:
239 if isinstance(self.data_struct._data_dict[key], torch.Tensor):
240 self.data_struct._data_dict[key] = self.data_struct._data_dict[key].cpu()
242 returned_struct = copy.deepcopy(self.data_struct)
243 for key in ["rec.topk", "rec.meanrank", "rec.score", "rec.items", "data.label", "rec.paths"]:
244 if key in self.data_struct:
245 del self.data_struct[key]
247 returned_struct.set("topk", self.topk)
249 return returned_struct
252class Collector_KG(Collector):
253 """This collector is used to collect the resource for evaluator in knowledge graph embedding models.
254 Specifically, it collects the predictions for the link prediction task, extending Collector from recommendation.
256 """
258 def __init__(self, config):
259 super().__init__(config)
260 self.register = Register_KG(config)
261 self.topk = self.config["topk_kg"]
264class ExplainableCollector(Collector):
265 """This collector is used to collect the resource for evaluator in explainable recommendation models.
266 It collects the KG paths and explanations for the recommendations made by the model and enables
267 path quality evaluation.
269 """
271 def __init__(self, config):
272 super().__init__(config)
273 self.register = Register(config)
275 def train_data_collect(self, train_data):
276 super().train_data_collect(train_data)
278 if self.register.need("data.max_path_type"):
279 self.data_struct.set("data.max_path_type", torch.arange(train_data.dataset.relation_num))
280 if self.register.need("data.node_degree"):
281 self.data_struct.set("data.node_degree", self.node_degree_dict(train_data))
282 if self.register.need("data.max_path_length"):
283 if hasattr(train_data.dataset, "token_sequence_length"):
284 # PEARLM or KGGLM
285 sampled_path_len = train_data.dataset.token_sequence_length - 2
286 self.data_struct.set("data.max_path_length", sampled_path_len)
287 else:
288 # PGPR or CAFE
289 self.data_struct.set(
290 "data.max_path_length", max([len(path) for path in self.config["path_constraint"]]) * 2 - 1
291 )
293 if self.register.need("data.rid2relation"):
294 self.data_struct.set("data.rid2relation", train_data.dataset.field2id_token["relation_id"])
296 def node_degree_dict(self, train_data):
297 # from pgpr knowledge graph
298 # https://github.com/giacoballoccu/rep-path-reasoning-recsys/blob/main/models/PGPR/knowledge_graph.py
299 aug_kg = train_data.dataset.ckg_dict_graph()
300 degrees = {}
301 for etype in aug_kg:
302 degrees[etype] = {}
303 for eid in aug_kg[etype]:
304 count = 0
305 for r in aug_kg[etype][eid]:
306 count += len(aug_kg[etype][eid][r])
307 degrees[etype][eid] = count
308 return degrees
310 def eval_batch_collect(
311 self,
312 explanations,
313 interaction,
314 positive_u: torch.Tensor,
315 positive_i: torch.Tensor,
316 ):
317 """Collect the evaluation resource from batched eval data and batched model output.
319 Args:
320 explanations (tuple): a tuple containing the scores and paths, where:
321 - scores (Torch.Tensor): the output tensor of model with the shape of `(N, )`
322 - paths (list): a list of quadruples representing the paths for each user.
323 interaction (Interaction): batched eval data.
324 positive_u (Torch.Tensor): the row index of positive items for each user.
325 positive_i (Torch.Tensor): the positive item id for each user.
326 """
327 scores, paths = explanations
328 super().eval_batch_collect(scores, interaction, positive_u, positive_i)
330 if self.register.need("rec.paths"):
331 self.data_struct.update_tensor("rec.paths", paths)