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

1# @Time : 2021/6/23 

2# @Author : Zihan Lin 

3# @Email : zhlin@ruc.edu.cn 

4 

5# UPDATE 

6# @Time : 2021/7/18 

7# @Author : Zhichao Feng 

8# @email : fzcbupt@gmail.com 

9 

10"""hopwise.evaluator.collector 

11################################################ 

12""" 

13 

14import copy 

15 

16import torch 

17 

18from hopwise.evaluator.register import Register, Register_KG 

19from hopwise.evaluator.utils import train_tsne 

20 

21 

22class DataStruct: 

23 def __init__(self): 

24 self._data_dict = {} 

25 

26 def __getitem__(self, name: str): 

27 return self._data_dict[name] 

28 

29 def __setitem__(self, name: str, value): 

30 self._data_dict[name] = value 

31 

32 def __delitem__(self, name: str): 

33 self._data_dict.pop(name) 

34 

35 def __contains__(self, key: str): 

36 return key in self._data_dict 

37 

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] 

42 

43 def set(self, name: str, value): 

44 self._data_dict[name] = value 

45 

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 

56 

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 

62 

63 

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. 

69 

70 This class is only used in Trainer. 

71 

72 """ 

73 

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"] 

81 

82 def train_data_collect(self, train_data): 

83 """Collect the evaluation resource from training data. 

84 

85 Args: 

86 train_data (AbstractDataLoader): the training dataloader which contains the training data. 

87 

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) 

106 

107 def eval_data_collect(self, eval_data): 

108 """Collect the evaluation resource from evaluation data, such as user and item features. 

109 

110 Args: 

111 eval_data (AbstractDataLoader): the evaluation dataloader which contains the evaluation data. 

112 

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) 

118 

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. 

121 

122 Args: 

123 scores(tensor): an ordered tensor, with size of `(N, )` 

124 

125 Returns: 

126 torch.Tensor: average_rank 

127 

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]]) 

132 

133 Reference: 

134 https://github.com/scipy/scipy/blob/v0.17.1/scipy/stats/stats.py#L5262-L5352 

135 

136 """ 

137 length, width = scores.shape 

138 true_tensor = torch.full((length, 1), True, dtype=torch.bool, device=self.device) 

139 

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 

144 

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) 

149 

150 return avg_rank 

151 

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. 

160 

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]) 

170 

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) 

175 

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) 

184 

185 if self.register.need("rec.meanrank"): 

186 desc_scores, desc_index = torch.sort(scores, dim=-1, descending=True) 

187 

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) 

192 

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) 

195 

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) 

200 

201 if self.register.need("rec.score"): 

202 self.data_struct.update_tensor("rec.score", scores) 

203 

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)) 

207 

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. 

210 

211 Args: 

212 model (nn.Module): the trained recommendation model. 

213 load_best_model (bool): whether to load the best model. 

214 """ 

215 

216 if self.config["tsne"] is not None: 

217 train_tsne(model, self.config["tsne"], load_best_model) 

218 

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. 

222 

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) 

229 

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)) 

233 

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() 

241 

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] 

246 

247 returned_struct.set("topk", self.topk) 

248 

249 return returned_struct 

250 

251 

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. 

255 

256 """ 

257 

258 def __init__(self, config): 

259 super().__init__(config) 

260 self.register = Register_KG(config) 

261 self.topk = self.config["topk_kg"] 

262 

263 

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. 

268 

269 """ 

270 

271 def __init__(self, config): 

272 super().__init__(config) 

273 self.register = Register(config) 

274 

275 def train_data_collect(self, train_data): 

276 super().train_data_collect(train_data) 

277 

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 ) 

292 

293 if self.register.need("data.rid2relation"): 

294 self.data_struct.set("data.rid2relation", train_data.dataset.field2id_token["relation_id"]) 

295 

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 

309 

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. 

318 

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) 

329 

330 if self.register.need("rec.paths"): 

331 self.data_struct.update_tensor("rec.paths", paths)