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

1# @Time : 2020/7/7 

2# @Author : Yupeng Hou 

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

4 

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 

9 

10"""hopwise.data.dataloader.abstract_dataloader 

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

12""" 

13 

14# ruff: noqa: PLW0602 PLW0603 

15 

16import copy 

17from logging import getLogger 

18 

19import torch 

20 

21from hopwise.data.interaction import Interaction 

22from hopwise.data.transform import construct_transform 

23from hopwise.utils import FeatureSource, FeatureType, InputType, ModelType 

24 

25start_iter = False 

26 

27 

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. 

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

46 

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 ) 

74 

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

78 

79 def update_config(self, config): 

80 """Update configure of dataloader, such as :attr:`batch_size`, :attr:`step` etc. 

81 

82 Args: 

83 config (Config): The new config of dataloader. 

84 """ 

85 self.config = config 

86 self._init_batch_size_and_step() 

87 

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. 

90 

91 Args: 

92 batch_size (int): the new batch_size of dataloader. 

93 """ 

94 self._batch_size = batch_size 

95 

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

99 

100 def __iter__(self): 

101 global start_iter 

102 start_iter = True 

103 res = super().__iter__() 

104 start_iter = False 

105 return res 

106 

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) 

112 

113 

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

118 

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

125 

126 def __init__(self, config, dataset, sampler, shuffle=True): 

127 self.logger = getLogger() 

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

129 

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

141 

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 

145 

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 

151 

152 self.neg_prefix = config["NEG_PREFIX"] 

153 self.neg_item_id = self.neg_prefix + self.iid_field 

154 

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

161 

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

164 

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 

191 

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 

199 

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 

209 

210 def get_model(self, model): 

211 self.model = model