Coverage for hopwise/data/dataloader/abstract_dataloader.py: 92%
113 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/10/22, 2020/9/23, 2022/7/6
7# @Author : Zhen Tian, Yupeng Hou, Yushuo Chen, Gaowei Zhang
8# @email : chenyuwuxinn@gmail.com, houyupeng@ruc.edu.cn, chenyushuo@ruc.edu.cn, zgw15630559577@163.com
10"""hopwise.data.dataloader.abstract_dataloader
11################################################
12"""
14# ruff: noqa: PLW0602 PLW0603
16import copy
17from logging import getLogger
19import torch
21from hopwise.data.interaction import Interaction
22from hopwise.data.transform import construct_transform
23from hopwise.utils import FeatureSource, FeatureType, InputType, ModelType
25start_iter = False
28class AbstractDataLoader(torch.utils.data.DataLoader):
29 """:class:`AbstractDataLoader` is an abstract object which would return a batch of data which is loaded by
30 :class:`~hopwise.data.interaction.Interaction` when it is iterated.
31 And it is also the ancestor of all other dataloader.
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``.
39 Attributes:
40 _dataset (Dataset): The dataset of this dataloader.
41 shuffle (bool): If ``True``, dataloader will shuffle before every epoch.
42 pr (int): Pointer of dataloader.
43 step (int): The increment of :attr:`pr` for each batch.
44 _batch_size (int): The max interaction number for all batch.
45 """
47 def __init__(self, config, dataset, sampler, shuffle=False):
48 self.shuffle = shuffle
49 self.config = config
50 self._dataset = dataset
51 self._sampler = sampler
52 self._batch_size = self.step = self.model = None
53 self._init_batch_size_and_step()
54 index_sampler = None
55 self.generator = torch.Generator()
56 self.generator.manual_seed(config["seed"])
57 self.transform = construct_transform(config)
58 self.is_sequential = config["MODEL_TYPE"] == ModelType.SEQUENTIAL
59 if not config["single_spec"]:
60 index_sampler = torch.utils.data.distributed.DistributedSampler(
61 list(range(self.sample_size)), shuffle=shuffle, drop_last=False
62 )
63 self.step = max(1, self.step // config["world_size"])
64 shuffle = False
65 super().__init__(
66 dataset=list(range(self.sample_size)),
67 batch_size=self.step,
68 collate_fn=self.collate_fn,
69 num_workers=config["worker"],
70 shuffle=shuffle,
71 sampler=index_sampler,
72 generator=self.generator,
73 )
75 def _init_batch_size_and_step(self):
76 """Initializing :attr:`step` and :attr:`batch_size`."""
77 raise NotImplementedError("Method [init_batch_size_and_step] should be implemented")
79 def update_config(self, config):
80 """Update configure of dataloader, such as :attr:`batch_size`, :attr:`step` etc.
82 Args:
83 config (Config): The new config of dataloader.
84 """
85 self.config = config
86 self._init_batch_size_and_step()
88 def set_batch_size(self, batch_size):
89 """Reset the batch_size of the dataloader, but it can't be called when dataloader is being iterated.
91 Args:
92 batch_size (int): the new batch_size of dataloader.
93 """
94 self._batch_size = batch_size
96 def collate_fn(self):
97 """Collect the sampled index, and apply neg_sampling or other methods to get the final data."""
98 raise NotImplementedError("Method [collate_fn] must be implemented.")
100 def __iter__(self):
101 global start_iter
102 start_iter = True
103 res = super().__iter__()
104 start_iter = False
105 return res
107 def __getattribute__(self, __name: str):
108 global start_iter
109 if not start_iter and __name == "dataset":
110 __name = "_dataset"
111 return super().__getattribute__(__name)
114class NegSampleDataLoader(AbstractDataLoader):
115 """:class:`NegSampleDataLoader` is an abstract class which can sample negative examples by ratio.
116 It has two neg-sampling method, the one is 1-by-1 neg-sampling (pair wise),
117 and the other is 1-by-multi neg-sampling (point wise).
119 Args:
120 config (Config): The config of dataloader.
121 dataset (Dataset): The dataset of dataloader.
122 sampler (Sampler): The sampler of dataloader.
123 shuffle (bool, optional): Whether the dataloader will be shuffle after a round. Defaults to ``False``.
124 """
126 def __init__(self, config, dataset, sampler, shuffle=True):
127 self.logger = getLogger()
128 super().__init__(config, dataset, sampler, shuffle=shuffle)
130 def _set_neg_sample_args(self, config, dataset, dl_format, neg_sample_args):
131 self.uid_field = dataset.uid_field
132 self.iid_field = dataset.iid_field
133 self.dl_format = dl_format
134 self.neg_sample_args = neg_sample_args
135 self.times = 1
136 if (
137 self.neg_sample_args["distribution"] in ["uniform", "popularity"]
138 and self.neg_sample_args["sample_num"] != "none"
139 ):
140 self.neg_sample_num = self.neg_sample_args["sample_num"]
142 if self.dl_format == InputType.POINTWISE:
143 self.times = 1 + self.neg_sample_num
144 self.sampling_func = self._neg_sample_by_point_wise_sampling
146 self.label_field = config["LABEL_FIELD"]
147 dataset.set_field_property(self.label_field, FeatureType.FLOAT, FeatureSource.INTERACTION, 1)
148 elif self.dl_format == InputType.PAIRWISE:
149 self.times = self.neg_sample_num
150 self.sampling_func = self._neg_sample_by_pair_wise_sampling
152 self.neg_prefix = config["NEG_PREFIX"]
153 self.neg_item_id = self.neg_prefix + self.iid_field
155 columns = [self.iid_field] if dataset.item_feat is None else dataset.item_feat.columns
156 for item_feat_col in columns:
157 neg_item_feat_col = self.neg_prefix + item_feat_col
158 dataset.copy_field_property(neg_item_feat_col, item_feat_col)
159 else:
160 raise ValueError(f"`neg sampling by` with dl_format [{self.dl_format}] not been implemented.")
162 elif self.neg_sample_args["distribution"] != "none" and self.neg_sample_args["sample_num"] != "none":
163 raise ValueError(f"`neg_sample_args` [{self.neg_sample_args['distribution']}] is not supported!")
165 def _neg_sampling(self, inter_feat):
166 if self.neg_sample_args.get("dynamic", False):
167 candidate_num = self.neg_sample_args["candidate_num"]
168 user_ids = inter_feat[self.uid_field].numpy()
169 item_ids = inter_feat[self.iid_field].numpy()
170 neg_candidate_ids = self._sampler.sample_by_user_ids(
171 user_ids, item_ids, self.neg_sample_num * candidate_num
172 )
173 self.model.eval()
174 interaction = copy.deepcopy(inter_feat).to(self.model.device)
175 interaction = interaction.repeat(self.neg_sample_num * candidate_num)
176 neg_item_feat = Interaction({self.iid_field: neg_candidate_ids.to(self.model.device)})
177 interaction.update(neg_item_feat)
178 scores = self.model.predict(interaction).reshape(candidate_num, -1)
179 indices = torch.max(scores, dim=0)[1].detach().cpu()
180 neg_candidate_ids = neg_candidate_ids.reshape(candidate_num, -1)
181 neg_item_ids = neg_candidate_ids[indices, [i for i in range(neg_candidate_ids.shape[1])]].view(-1)
182 self.model.train()
183 return self.sampling_func(inter_feat, neg_item_ids)
184 elif self.neg_sample_args["distribution"] != "none" and self.neg_sample_args["sample_num"] != "none":
185 user_ids = inter_feat[self.uid_field].numpy()
186 item_ids = inter_feat[self.iid_field].numpy()
187 neg_item_ids = self._sampler.sample_by_user_ids(user_ids, item_ids, self.neg_sample_num)
188 return self.sampling_func(inter_feat, neg_item_ids)
189 else:
190 return inter_feat
192 def _neg_sample_by_pair_wise_sampling(self, inter_feat, neg_item_ids):
193 inter_feat = inter_feat.repeat(self.times)
194 neg_item_feat = Interaction({self.iid_field: neg_item_ids})
195 neg_item_feat = self._dataset.join(neg_item_feat)
196 neg_item_feat.add_prefix(self.neg_prefix)
197 inter_feat.update(neg_item_feat)
198 return inter_feat
200 def _neg_sample_by_point_wise_sampling(self, inter_feat, neg_item_ids):
201 pos_inter_num = len(inter_feat)
202 new_data = inter_feat.repeat(self.times)
203 new_data[self.iid_field][pos_inter_num:] = neg_item_ids
204 new_data = self._dataset.join(new_data)
205 labels = torch.zeros(pos_inter_num * self.times)
206 labels[:pos_inter_num] = 1.0
207 new_data.update(Interaction({self.label_field: labels}))
208 return new_data
210 def get_model(self, model):
211 self.model = model