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
« 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
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
10"""hopwise.data.dataloader.general_dataloader
11################################################
12"""
14from logging import getLogger
16import numpy as np
17import torch
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
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.
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 """
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)
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)
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)
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)
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.
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 """
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)
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)
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)
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)
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)
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()
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
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
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 """
181 def __init__(self, config, dataset, sampler, shuffle=False):
182 self.logger = getLogger()
184 if shuffle:
185 self.logger.warning("FullSortEvalDataLoader can't shuffle")
186 shuffle = False
187 super().__init__(config, dataset, sampler, shuffle=shuffle)
189 def check_sequential(self, config):
190 self.is_sequential = config["MODEL_TYPE"] == ModelType.SEQUENTIAL
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])
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)
201 feat.sort(by=source_field, ascending=True)
202 last_source = None
203 positives = set()
204 used_ids = sampler.used_ids
206 if extra_fields is None:
207 extra_fields = []
208 elif not isinstance(extra_fields, list):
209 extra_fields = [extra_fields]
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)
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)
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)
250 def update_config(self, config):
251 super().update_config(config)
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])
258 history = self._sample2history[source_list]
259 positives = self._sample2positives[source_list]
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))
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))
267 return source_df, (history_source, history_target), positive_source, positive_target
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]
280 return transformed_interaction, None, positive_u, positive_i
283class FullSortRecEvalDataLoader(FullSortEvalDataLoader):
284 """:class:`FullSortRecEvalDataLoader` is a dataloader for full-sort evaluation for the recommendation (Rec) task.
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 """
293 def __init__(self, config, dataset, sampler, shuffle=False):
294 self.check_sequential(config)
296 # needed for TPRec
297 if hasattr(dataset, "temporal_weights"):
298 self.temporal_weights = dataset.temporal_weights
300 self.uid_field = dataset.uid_field
301 self.iid_field = dataset.iid_field
302 self._source_field = self.uid_field
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
311 super().__init__(config, dataset, sampler, shuffle=shuffle)
314class FullSortLPEvalDataLoader(FullSortEvalDataLoader):
315 """:class:`FullSortLPEvalDataLoader` is a dataloader for full-sort evaluation for the link prediction (LP) task.
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 """
324 def __init__(self, config, dataset, sampler, shuffle=False):
325 self.check_sequential(config)
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
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
345 super().__init__(config, dataset, sampler, shuffle=shuffle)