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
« 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/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
10# UPDATE
11# @Time : 2025
12# @Author : Giacomo Medda
13# @Email : giacomo.medda@unica.it
16"""hopwise.data.dataloader.knowledge_dataloader
17################################################
18"""
20from logging import getLogger
22import numpy as np
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
30class KGDataLoader(AbstractDataLoader):
31 """:class:`KGDataLoader` is a dataloader which would return the triplets with negative examples
32 in a knowledge graph.
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``.
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 """
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")
51 self.neg_sample_num = 1
53 self.neg_prefix = config["NEG_PREFIX"]
54 self.hid_field = dataset.head_entity_field
55 self.tid_field = dataset.tail_entity_field
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)
61 self.sample_size = len(dataset.kg_feat)
62 super().__init__(config, dataset, sampler, shuffle=shuffle)
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)
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
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`.
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``.
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`
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
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)
110 # using kg_sampler
111 self.kg_dataloader = KGDataLoader(config, dataset, kg_sampler, shuffle=True)
113 self.shuffle = False
114 self.state = None
115 self.dataset = self._dataset = dataset
116 self.kg_iter, self.gen_iter = None, None
118 def update_config(self, config):
119 self.general_dataloader.update_config(config)
120 self.kg_dataloader.update_config(config)
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
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
147 def __len__(self):
148 if self.state == KGDataLoaderState.KG:
149 return len(self.kg_dataloader)
150 else:
151 return len(self.general_dataloader)
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
159 The state of :class:`KnowledgeBasedDataLoader` would affect the result of _next_batch_data().
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
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)
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)
176 if self.general_dataloader.shuffle:
177 self.general_dataloader.sampler.set_epoch(epoch_seed)
180class KnowledgePathEvalDataLoader(FullSortRecEvalDataLoader):
181 def __init__(self, config, dataset, sampler, shuffle=False):
182 super().__init__(config, dataset, sampler, shuffle)
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)
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)
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