Coverage for hopwise/data/dataloader/knowledge_dataloader.py: 83%

92 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/18, 2020/9/21, 2020/8/31 

7# @Author : Zhen Tian, Yupeng Hou, Yushuo Chen, Kaiyuan Li 

8# @email : chenyuwuxinn@gmail.com, houyupeng@ruc.edu.cn, chenyushuo@ruc.edu.cn, tsotfsk@outlook.com 

9 

10# UPDATE 

11# @Time : 2025 

12# @Author : Giacomo Medda 

13# @Email : giacomo.medda@unica.it 

14 

15 

16"""hopwise.data.dataloader.knowledge_dataloader 

17################################################ 

18""" 

19 

20from logging import getLogger 

21 

22import numpy as np 

23 

24from hopwise.data.dataloader.abstract_dataloader import AbstractDataLoader 

25from hopwise.data.dataloader.general_dataloader import FullSortRecEvalDataLoader, TrainDataLoader 

26from hopwise.data.interaction import Interaction 

27from hopwise.utils import KGDataLoaderState, PathLanguageModelingTokenType 

28 

29 

30class KGDataLoader(AbstractDataLoader): 

31 """:class:`KGDataLoader` is a dataloader which would return the triplets with negative examples 

32 in a knowledge graph. 

33 

34 Args: 

35 config (Config): The config of dataloader. 

36 dataset (Dataset): The dataset of dataloader. 

37 sampler (KGSampler): The knowledge graph sampler of dataloader. 

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

39 

40 Attributes: 

41 shuffle (bool): Whether the dataloader will be shuffle after a round. 

42 However, in :class:`KGDataLoader`, it's guaranteed to be ``True``. 

43 """ 

44 

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

46 self.logger = getLogger() 

47 if shuffle is False: 

48 shuffle = True 

49 self.logger.warning("kg based dataloader must shuffle the data") 

50 

51 self.neg_sample_num = 1 

52 

53 self.neg_prefix = config["NEG_PREFIX"] 

54 self.hid_field = dataset.head_entity_field 

55 self.tid_field = dataset.tail_entity_field 

56 

57 # kg negative cols 

58 self.neg_tid_field = self.neg_prefix + self.tid_field 

59 dataset.copy_field_property(self.neg_tid_field, self.tid_field) 

60 

61 self.sample_size = len(dataset.kg_feat) 

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

63 

64 def _init_batch_size_and_step(self): 

65 batch_size = self.config["train_batch_size"] 

66 self.step = batch_size 

67 self.set_batch_size(batch_size) 

68 

69 def collate_fn(self, index): 

70 index = np.array(index) 

71 cur_data = self._dataset.kg_feat[index] 

72 head_ids = cur_data[self.hid_field].numpy() 

73 neg_tail_ids = self._sampler.sample_by_entity_ids(head_ids, self.neg_sample_num) 

74 cur_data.update(Interaction({self.neg_tid_field: neg_tail_ids})) 

75 return cur_data 

76 

77 

78class KnowledgeBasedDataLoader: 

79 """:class:`KnowledgeBasedDataLoader` is used for knowledge based model. 

80 It has three states, which is saved in :attr:`state`. 

81 In different states, :meth:`~_next_batch_data` will return different :class:`~hopwise.data.interaction.Interaction`. 

82 Detailed, please see :attr:`~state`. 

83 

84 Args: 

85 config (Config): The config of dataloader. 

86 dataset (Dataset): The dataset of dataloader. 

87 sampler (Sampler): The sampler of dataloader. 

88 kg_sampler (KGSampler): The knowledge graph sampler of dataloader. 

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

90 

91 Attributes: 

92 state (KGDataLoaderState): 

93 This dataloader has three states: 

94 - :obj:`~hopwise.utils.enum_type.KGDataLoaderState.RS` 

95 - :obj:`~hopwise.utils.enum_type.KGDataLoaderState.KG` 

96 - :obj:`~hopwise.utils.enum_type.KGDataLoaderState.RSKG` 

97 

98 In the first state, this dataloader would only return the user-item interaction. 

99 In the second state, this dataloader would only return the triplets with negative 

100 examples in a knowledge graph. 

101 In the last state, this dataloader would return both knowledge graph information 

102 and user-item interaction information. 

103 """ # noqa: E501 

104 

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

106 self.logger = getLogger() 

107 # using sampler 

108 self.general_dataloader = TrainDataLoader(config, dataset, sampler, shuffle=shuffle) 

109 

110 # using kg_sampler 

111 self.kg_dataloader = KGDataLoader(config, dataset, kg_sampler, shuffle=True) 

112 

113 self.shuffle = False 

114 self.state = None 

115 self.dataset = self._dataset = dataset 

116 self.kg_iter, self.gen_iter = None, None 

117 

118 def update_config(self, config): 

119 self.general_dataloader.update_config(config) 

120 self.kg_dataloader.update_config(config) 

121 

122 def __iter__(self): 

123 if self.state is None: 

124 raise ValueError( 

125 "The dataloader's state must be set when using the kg based dataloader, " 

126 "you should call set_mode() before __iter__()" 

127 ) 

128 if self.state == KGDataLoaderState.KG: 

129 return self.kg_dataloader.__iter__() 

130 elif self.state == KGDataLoaderState.RS: 

131 return self.general_dataloader.__iter__() 

132 elif self.state == KGDataLoaderState.RSKG: 

133 self.kg_iter = self.kg_dataloader.__iter__() 

134 self.gen_iter = self.general_dataloader.__iter__() 

135 return self 

136 

137 def __next__(self): 

138 try: 

139 kg_data = next(self.kg_iter) 

140 except StopIteration: 

141 self.kg_iter = self.kg_dataloader.__iter__() 

142 kg_data = next(self.kg_iter) 

143 recdata = next(self.gen_iter) 

144 recdata.update(kg_data) 

145 return recdata 

146 

147 def __len__(self): 

148 if self.state == KGDataLoaderState.KG: 

149 return len(self.kg_dataloader) 

150 else: 

151 return len(self.general_dataloader) 

152 

153 def set_mode(self, state): 

154 """Set the mode of :class:`KnowledgeBasedDataLoader`, it can be set to three states: 

155 - KGDataLoaderState.RS 

156 - KGDataLoaderState.KG 

157 - KGDataLoaderState.RSKG 

158 

159 The state of :class:`KnowledgeBasedDataLoader` would affect the result of _next_batch_data(). 

160 

161 Args: 

162 state (KGDataLoaderState): the state of :class:`KnowledgeBasedDataLoader`. 

163 """ 

164 if state not in set(KGDataLoaderState): 

165 raise NotImplementedError(f"Kg data loader has no state named [{self.state}].") 

166 self.state = state 

167 

168 def get_model(self, model): 

169 """Let the general_dataloader get the model, used for dynamic sampling.""" 

170 self.general_dataloader.get_model(model) 

171 

172 def knowledge_shuffle(self, epoch_seed): 

173 """Reset the seed to ensure that each subprocess generates the same index squence.""" 

174 self.kg_dataloader.sampler.set_epoch(epoch_seed) 

175 

176 if self.general_dataloader.shuffle: 

177 self.general_dataloader.sampler.set_epoch(epoch_seed) 

178 

179 

180class KnowledgePathEvalDataLoader(FullSortRecEvalDataLoader): 

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

182 super().__init__(config, dataset, sampler, shuffle) 

183 

184 user_df = self.user_df[self.uid_field] 

185 ui_relation = dataset.field2token_id[dataset.relation_field][dataset.ui_relation] 

186 inference_path_dataset = [ 

187 dataset.path_token_separator.join( 

188 [ 

189 dataset.tokenizer.bos_token, 

190 PathLanguageModelingTokenType.USER.token + str(uid.item()), 

191 PathLanguageModelingTokenType.RELATION.token + str(ui_relation), 

192 ] 

193 ) 

194 for uid in user_df 

195 ] 

196 inference_tokenized_dataset = dataset.tokenizer( 

197 inference_path_dataset, return_tensors="pt", add_special_tokens=False 

198 ) 

199 self.inference_tokenized_dataset = Interaction(inference_tokenized_dataset.data) 

200 

201 def _init_batch_size_and_step(self): 

202 batch_size = self.config["eval_batch_size"] 

203 self.step = batch_size 

204 self.set_batch_size(batch_size) 

205 

206 def collate_fn(self, index): 

207 _, history_index, positive_u, positive_i = super().collate_fn(index) 

208 return self.inference_tokenized_dataset[index], history_index, positive_u, positive_i