Coverage for hopwise/data/dataloader/general_dataloader.py: 87%

207 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1# @Time : 2020/7/7 

2# @Author : Yupeng Hou 

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

4 

5# UPDATE 

6# @Time : 2022/7/8, 2020/9/9, 2020/9/29, 2021/7/15, 2022/7/6 

7# @Author : Zhen Tian, Yupeng Hou, Yushuo Chen, Xingyu Pan, Gaowei Zhang 

8# @email : chenyuwuxinn@gmail.com, houyupeng@ruc.edu.cn, chenyushuo@ruc.edu.cn, xy_pan@foxmail.com, zgw15630559577@163.com # noqa: E501 

9 

10"""hopwise.data.dataloader.general_dataloader 

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

12""" 

13 

14from logging import getLogger 

15 

16import numpy as np 

17import torch 

18 

19from hopwise.data.dataloader.abstract_dataloader import ( 

20 AbstractDataLoader, 

21 NegSampleDataLoader, 

22) 

23from hopwise.data.interaction import Interaction, cat_interactions 

24from hopwise.utils import InputType, ModelType 

25 

26 

27class TrainDataLoader(NegSampleDataLoader): 

28 """:class:`TrainDataLoader` is a dataloader for training. 

29 It can generate negative interaction when :attr:`training_neg_sample_num` is not zero. 

30 For the result of every batch, we permit that every positive interaction and its negative interaction 

31 must be in the same batch. 

32 

33 Args: 

34 config (Config): The config of dataloader. 

35 dataset (Dataset): The dataset of dataloader. 

36 sampler (Sampler): The sampler of dataloader. 

37 shuffle (bool, optional): Whether the dataloader will be shuffle after a round. Defaults to ``False``. 

38 """ 

39 

40 def __init__(self, config, dataset, sampler, shuffle=False): 

41 self.logger = getLogger() 

42 self._set_neg_sample_args(config, dataset, config["MODEL_INPUT_TYPE"], config["train_neg_sample_args"]) 

43 self.sample_size = len(dataset) 

44 super().__init__(config, dataset, sampler, shuffle=shuffle) 

45 

46 def _init_batch_size_and_step(self): 

47 batch_size = self.config["train_batch_size"] 

48 if self.neg_sample_args["distribution"] != "none": 

49 batch_num = max(batch_size // self.times, 1) 

50 new_batch_size = batch_num * self.times 

51 self.step = batch_num 

52 self.set_batch_size(new_batch_size) 

53 else: 

54 self.step = batch_size 

55 self.set_batch_size(batch_size) 

56 

57 def update_config(self, config): 

58 self._set_neg_sample_args( 

59 config, 

60 self._dataset, 

61 config["MODEL_INPUT_TYPE"], 

62 config["train_neg_sample_args"], 

63 ) 

64 super().update_config(config) 

65 

66 def collate_fn(self, index): 

67 index = np.array(index) 

68 data = self._dataset[index] 

69 transformed_data = self.transform(self._dataset, data) 

70 return self._neg_sampling(transformed_data) 

71 

72 

73class NegSampleEvalDataLoader(NegSampleDataLoader): 

74 """:class:`NegSampleEvalDataLoader` is a dataloader for neg-sampling evaluation. 

75 It is similar to :class:`TrainDataLoader` which can generate negative items, 

76 and this dataloader also permits that all the interactions corresponding to each user are in the same batch 

77 and positive interactions are before negative interactions. 

78 

79 Args: 

80 config (Config): The config of dataloader. 

81 dataset (Dataset): The dataset of dataloader. 

82 sampler (Sampler): The sampler of dataloader. 

83 shuffle (bool, optional): Whether the dataloader will be shuffle after a round. Defaults to ``False``. 

84 """ 

85 

86 def __init__(self, config, dataset, sampler, shuffle=False): 

87 self.logger = getLogger() 

88 phase = sampler.phase if sampler is not None else "test" 

89 self._set_neg_sample_args(config, dataset, InputType.POINTWISE, config[f"{phase}_neg_sample_args"]) 

90 if self.neg_sample_args["distribution"] != "none" and self.neg_sample_args["sample_num"] != "none": 

91 user_num = dataset.user_num 

92 dataset.sort(by=dataset.uid_field, ascending=True) 

93 self.uid_list = [] 

94 start, end = dict(), dict() 

95 for i, uid in enumerate(dataset.inter_feat[dataset.uid_field].numpy()): 

96 if uid not in start: 

97 self.uid_list.append(uid) 

98 start[uid] = i 

99 end[uid] = i 

100 self.uid2index = np.array([None] * user_num) 

101 self.uid2items_num = np.zeros(user_num, dtype=np.int64) 

102 for uid in self.uid_list: 

103 self.uid2index[uid] = slice(start[uid], end[uid] + 1) 

104 self.uid2items_num[uid] = end[uid] - start[uid] + 1 

105 self.uid_list = np.array(self.uid_list) 

106 self.sample_size = len(self.uid_list) 

107 else: 

108 self.sample_size = len(dataset) 

109 if shuffle: 

110 self.logger.warning("NegSampleEvalDataLoader can't shuffle") 

111 shuffle = False 

112 super().__init__(config, dataset, sampler, shuffle=shuffle) 

113 

114 def _init_batch_size_and_step(self): 

115 batch_size = self.config["eval_batch_size"] 

116 if self.neg_sample_args["distribution"] != "none" and self.neg_sample_args["sample_num"] != "none": 

117 inters_num = sorted(self.uid2items_num * self.times, reverse=True) 

118 batch_num = 1 

119 new_batch_size = inters_num[0] 

120 for i in range(1, len(inters_num)): 

121 if new_batch_size + inters_num[i] > batch_size: 

122 break 

123 batch_num = i + 1 

124 new_batch_size += inters_num[i] 

125 self.step = batch_num 

126 self.set_batch_size(new_batch_size) 

127 else: 

128 self.step = batch_size 

129 self.set_batch_size(batch_size) 

130 

131 def update_config(self, config): 

132 phase = self._sampler.phase if self._sampler.phase is not None else "test" 

133 self._set_neg_sample_args( 

134 config, 

135 self._dataset, 

136 InputType.POINTWISE, 

137 config[f"{phase}_neg_sample_args"], 

138 ) 

139 super().update_config(config) 

140 

141 def collate_fn(self, index): 

142 index = np.array(index) 

143 if self.neg_sample_args["distribution"] != "none" and self.neg_sample_args["sample_num"] != "none": 

144 uid_list = self.uid_list[index] 

145 data_list = [] 

146 idx_list = [] 

147 positive_u = [] 

148 positive_i = torch.tensor([], dtype=torch.int64) 

149 

150 for idx, uid in enumerate(uid_list): 

151 index = self.uid2index[uid] 

152 transformed_data = self.transform(self._dataset, self._dataset[index]) 

153 data_list.append(self._neg_sampling(transformed_data)) 

154 idx_list += [idx for i in range(self.uid2items_num[uid] * self.times)] 

155 positive_u += [idx for i in range(self.uid2items_num[uid])] 

156 positive_i = torch.cat((positive_i, self._dataset[index][self.iid_field]), 0) 

157 

158 cur_data = cat_interactions(data_list) 

159 idx_list = torch.from_numpy(np.array(idx_list)).long() 

160 positive_u = torch.from_numpy(np.array(positive_u)).long() 

161 

162 return cur_data, idx_list, positive_u, positive_i 

163 else: 

164 data = self._dataset[index] 

165 transformed_data = self.transform(self._dataset, data) 

166 cur_data = self._neg_sampling(transformed_data) 

167 return cur_data, None, None, None 

168 

169 

170class FullSortEvalDataLoader(AbstractDataLoader): 

171 """:class:`FullSortEvalDataLoader` is a dataloader for full-sort evaluation. In order to speed up calculation, 

172 this dataloader would only return the data samples with positives, not negatives 

173 

174 Args: 

175 config (Config): The config of dataloader. 

176 dataset (Dataset): The dataset of dataloader. 

177 sampler (Sampler): The sampler of dataloader. 

178 shuffle (bool, optional): Whether the dataloader will be shuffle after a round. Defaults to ``False``. 

179 """ 

180 

181 def __init__(self, config, dataset, sampler, shuffle=False): 

182 self.logger = getLogger() 

183 

184 if shuffle: 

185 self.logger.warning("FullSortEvalDataLoader can't shuffle") 

186 shuffle = False 

187 super().__init__(config, dataset, sampler, shuffle=shuffle) 

188 

189 def check_sequential(self, config): 

190 self.is_sequential = config["MODEL_TYPE"] == ModelType.SEQUENTIAL 

191 

192 def _build_positive_samples(self, dataset, sampler, feat, target_field, extra_fields=None): 

193 source_field = self._source_field 

194 source_num = len(dataset.field2id_token[source_field]) 

195 

196 self._source_list = [] 

197 self._sample2positive_num = np.zeros(source_num, dtype=np.int64) 

198 self._sample2positives = np.array([None] * source_num) 

199 self._sample2history = np.array([None] * source_num) 

200 

201 feat.sort(by=source_field, ascending=True) 

202 last_source = None 

203 positives = set() 

204 used_ids = sampler.used_ids 

205 

206 if extra_fields is None: 

207 extra_fields = [] 

208 elif not isinstance(extra_fields, list): 

209 extra_fields = [extra_fields] 

210 

211 feat_extra_fields = [feat[field].numpy() for field in extra_fields] 

212 extra_fields_list = {field: [] for field in extra_fields} 

213 for source, target, *add_fields_values in zip( 

214 feat[source_field].numpy(), feat[target_field].numpy(), *feat_extra_fields 

215 ): 

216 if source != last_source: 

217 self._set_source_property(last_source, used_ids[last_source], positives) 

218 last_source = source 

219 self._source_list.append(source) 

220 positives = set() 

221 for field, value in zip(extra_fields, add_fields_values): 

222 extra_fields_list[field].append(value) 

223 positives.add(target) 

224 self._set_source_property(last_source, used_ids[last_source], positives) 

225 self._source_list = torch.tensor(self._source_list, dtype=torch.int64) 

226 for field in extra_fields: 

227 extra_fields_list[field] = torch.tensor(extra_fields_list[field], dtype=torch.int64) 

228 self._source_df = dataset.join(Interaction({source_field: self._source_list, **extra_fields_list})) 

229 self.sample_size = len(self._source_df) if not self.is_sequential else len(dataset) 

230 

231 def _set_source_property(self, source, used_ids, positives): 

232 if source is None: 

233 return 

234 history = used_ids - positives 

235 self._sample2positives[source] = torch.tensor(list(positives), dtype=torch.int64) 

236 self._sample2positive_num[source] = len(positives) 

237 self._sample2history[source] = torch.tensor(list(history), dtype=torch.int64) 

238 

239 def _init_batch_size_and_step(self): 

240 batch_size = self.config["eval_batch_size"] 

241 if not self.is_sequential: 

242 batch_num = max(batch_size // self._dataset.item_num, 1) 

243 new_batch_size = batch_num * self._dataset.item_num 

244 self.step = batch_num 

245 self.set_batch_size(new_batch_size) 

246 else: 

247 self.step = batch_size 

248 self.set_batch_size(batch_size) 

249 

250 def update_config(self, config): 

251 super().update_config(config) 

252 

253 def _not_sequential_collate_fn(self, index, source_field): 

254 index = np.array(index) 

255 source_df = self._source_df[index] 

256 source_list = list(source_df[source_field]) 

257 

258 history = self._sample2history[source_list] 

259 positives = self._sample2positives[source_list] 

260 

261 history_source = torch.cat([torch.full_like(hist_iid, i) for i, hist_iid in enumerate(history)]) 

262 history_target = torch.cat(list(history)) 

263 

264 positive_source = torch.cat([torch.full_like(pos_iid, i) for i, pos_iid in enumerate(positives)]) 

265 positive_target = torch.cat(list(positives)) 

266 

267 return source_df, (history_source, history_target), positive_source, positive_target 

268 

269 def collate_fn(self, index): 

270 index = np.array(index) 

271 if not self.is_sequential: 

272 return self._not_sequential_collate_fn(index, self._source_field) 

273 else: 

274 interaction = self._dataset[index] 

275 transformed_interaction = self.transform(self._dataset, interaction) 

276 inter_num = len(transformed_interaction) 

277 positive_u = torch.arange(inter_num) 

278 positive_i = transformed_interaction[self.iid_field] 

279 

280 return transformed_interaction, None, positive_u, positive_i 

281 

282 

283class FullSortRecEvalDataLoader(FullSortEvalDataLoader): 

284 """:class:`FullSortRecEvalDataLoader` is a dataloader for full-sort evaluation for the recommendation (Rec) task. 

285 

286 Args: 

287 config (Config): The config of dataloader. 

288 dataset (Dataset): The dataset of dataloader. 

289 sampler (Sampler): The sampler of dataloader. 

290 shuffle (bool, optional): Whether the dataloader will be shuffle after a round. Defaults to ``False``. 

291 """ 

292 

293 def __init__(self, config, dataset, sampler, shuffle=False): 

294 self.check_sequential(config) 

295 

296 # needed for TPRec 

297 if hasattr(dataset, "temporal_weights"): 

298 self.temporal_weights = dataset.temporal_weights 

299 

300 self.uid_field = dataset.uid_field 

301 self.iid_field = dataset.iid_field 

302 self._source_field = self.uid_field 

303 

304 self._build_positive_samples(dataset, sampler, dataset.inter_feat, self.iid_field) 

305 self.uid2items_num = self._sample2positive_num 

306 self.uid2positive_item = self._sample2positives 

307 self.uid2history_item = self._sample2history 

308 self.uid_list = self._source_list 

309 self.user_df = self._source_df 

310 

311 super().__init__(config, dataset, sampler, shuffle=shuffle) 

312 

313 

314class FullSortLPEvalDataLoader(FullSortEvalDataLoader): 

315 """:class:`FullSortLPEvalDataLoader` is a dataloader for full-sort evaluation for the link prediction (LP) task. 

316 

317 Args: 

318 config (Config): The config of dataloader. 

319 dataset (Dataset): The dataset of dataloader. 

320 sampler (Sampler): The sampler of dataloader. 

321 shuffle (bool, optional): Whether the dataloader will be shuffle after a round. Defaults to ``False``. 

322 """ 

323 

324 def __init__(self, config, dataset, sampler, shuffle=False): 

325 self.check_sequential(config) 

326 

327 self.head_entity_field = dataset.head_entity_field 

328 self.relation_field = dataset.relation_field 

329 self.tail_entity_field = dataset.tail_entity_field 

330 self._source_field = self.head_entity_field 

331 

332 self._build_positive_samples( 

333 dataset, 

334 sampler, 

335 dataset.kg_feat, 

336 self.tail_entity_field, 

337 extra_fields=[self.relation_field], 

338 ) 

339 self.head2tails_num = self._sample2positive_num 

340 self.head2positive_tail = self._sample2positives 

341 self.head2history_tail = self._sample2history 

342 self.head_list = self._source_list 

343 self.kg_df = self._source_df 

344 

345 super().__init__(config, dataset, sampler, shuffle=shuffle)