Coverage for hopwise/sampler/sampler.py: 78%
209 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# @Author : Yupeng Hou
2# @Email : houyupeng@ruc.edu.cn
3# @File : sampler.py
5# UPDATE
6# @Time : 2021/7/23, 2020/8/31, 2020/10/6, 2020/9/18, 2021/3/19
7# @Author : Xingyu Pan, Kaiyuan Li, Yupeng Hou, Yushuo Chen, Zhichao Feng
8# @email : xy_pan@foxmail.com, tsotfsk@outlook.com, houyupeng@ruc.edu.cn, chenyushuo@ruc.edu.cn, fzcbupt@gmail.com
10"""hopwise.sampler
11########################
12"""
14import copy
15from collections import Counter
17import numpy as np
18import torch
21class AbstractSampler:
22 """:class:`AbstractSampler` is a abstract class, all sampler should inherit from it. This sampler supports
23 returning a certain number of random value_ids according to the input key_id, and it also supports
24 to prohibit certain key-value pairs by setting used_ids.
26 Args:
27 distribution (str): The string of distribution, which is used for subclass.
29 Attributes:
30 used_ids (numpy.ndarray): The result of :meth:`get_used_ids`.
31 """
33 def __init__(self, distribution, alpha):
34 self.distribution = ""
35 self.alpha = alpha
36 self.set_distribution(distribution)
37 self.used_ids = self.get_used_ids()
39 def set_distribution(self, distribution):
40 """Set the distribution of sampler.
42 Args:
43 distribution (str): Distribution of the negative items.
44 """
45 self.distribution = distribution
46 if distribution == "popularity":
47 self._build_alias_table()
49 def _uni_sampling(self, sample_num):
50 """Sample [sample_num] items in the uniform distribution.
52 Args:
53 sample_num (int): the number of samples.
55 Returns:
56 sample_list (np.array): a list of samples.
57 """
58 raise NotImplementedError("Method [_uni_sampling] should be implemented")
60 def _get_candidates_list(self):
61 """Get sample candidates list for _pop_sampling()
63 Returns:
64 candidates_list (list): a list of candidates id.
65 """
66 raise NotImplementedError("Method [_get_candidates_list] should be implemented")
68 def _build_alias_table(self):
69 """Build alias table for popularity_biased sampling."""
70 candidates_list = self._get_candidates_list()
71 self.prob = dict(Counter(candidates_list))
72 self.alias = self.prob.copy()
73 large_q = []
74 small_q = []
75 for i in self.prob:
76 self.alias[i] = -1
77 self.prob[i] = self.prob[i] / len(candidates_list)
78 self.prob[i] = pow(self.prob[i], self.alpha)
79 normalize_count = sum(self.prob.values())
80 for i in self.prob:
81 self.prob[i] = self.prob[i] / normalize_count * len(self.prob)
82 if self.prob[i] > 1:
83 large_q.append(i)
84 elif self.prob[i] < 1:
85 small_q.append(i)
86 while len(large_q) != 0 and len(small_q) != 0:
87 lq_el = large_q.pop(0)
88 sq_el = small_q.pop(0)
89 self.alias[sq_el] = lq_el
90 self.prob[lq_el] = self.prob[lq_el] - (1 - self.prob[sq_el])
91 if self.prob[lq_el] < 1:
92 small_q.append(lq_el)
93 elif self.prob[lq_el] > 1:
94 large_q.append(lq_el)
96 def _pop_sampling(self, sample_num):
97 """Sample [sample_num] items in the popularity-biased distribution.
99 Args:
100 sample_num (int): the number of samples.
102 Returns:
103 sample_list (np.array): a list of samples.
104 """
105 keys = list(self.prob.keys())
106 random_index_list = np.random.randint(0, len(keys), sample_num)
107 random_prob_list = np.random.random(sample_num)
108 final_random_list = []
110 for idx, prob in zip(random_index_list, random_prob_list):
111 if self.prob[keys[idx]] > prob:
112 final_random_list.append(keys[idx])
113 else:
114 final_random_list.append(self.alias[keys[idx]])
116 return np.array(final_random_list)
118 def sampling(self, sample_num):
119 """Sampling [sample_num] item_ids.
121 Args:
122 sample_num (int): the number of samples.
124 Returns:
125 sample_list (np.array): a list of samples and the len is [sample_num].
126 """
127 if self.distribution == "uniform":
128 return self._uni_sampling(sample_num)
129 elif self.distribution == "popularity":
130 return self._pop_sampling(sample_num)
131 else:
132 raise NotImplementedError(f"The sampling distribution [{self.distribution}] is not implemented.")
134 def get_used_ids(self):
135 """Returns:
136 numpy.ndarray: Used ids. Index is key_id, and element is a set of value_ids.
137 """
138 raise NotImplementedError("Method [get_used_ids] should be implemented")
140 def sample_by_key_ids(self, key_ids, num):
141 """Sampling by key_ids.
143 Args:
144 key_ids (numpy.ndarray or list): Input key_ids.
145 num (int): Number of sampled value_ids for each key_id.
147 Returns:
148 torch.tensor: Sampled value_ids.
149 value_ids[0], value_ids[len(key_ids)], value_ids[len(key_ids) * 2], ..., value_id[len(key_ids) * (num - 1)]
150 is sampled for key_ids[0];
151 value_ids[1], value_ids[len(key_ids) + 1], value_ids[len(key_ids) * 2 + 1], ...,
152 value_id[len(key_ids) * (num - 1) + 1] is sampled for key_ids[1]; ...; and so on.
153 """
154 key_ids = np.array(key_ids)
155 key_num = len(key_ids)
156 total_num = key_num * num
157 if (key_ids == key_ids[0]).all():
158 key_id = key_ids[0]
159 used = np.array(list(self.used_ids[key_id]))
160 value_ids = self.sampling(total_num)
161 check_list = np.arange(total_num)[np.isin(value_ids, used)]
162 while len(check_list) > 0:
163 value_ids[check_list] = value = self.sampling(len(check_list))
164 mask = np.isin(value, used)
165 check_list = check_list[mask]
166 else:
167 value_ids = np.zeros(total_num, dtype=np.int64)
168 check_list = np.arange(total_num)
169 key_ids = np.tile(key_ids, num)
170 while len(check_list) > 0:
171 value_ids[check_list] = self.sampling(len(check_list))
172 check_list = np.array(
173 [
174 i
175 for i, used, v in zip(
176 check_list,
177 self.used_ids[key_ids[check_list]],
178 value_ids[check_list],
179 )
180 if v in used
181 ]
182 )
183 return torch.tensor(value_ids, dtype=torch.long)
186class Sampler(AbstractSampler):
187 """:class:`Sampler` is used to sample negative items for each input user. In order to avoid positive items
188 in train-phase to be sampled in valid-phase, and positive items in train-phase or valid-phase to be sampled
189 in test-phase, we need to input the datasets of all phases for pre-processing. And, before using this sampler,
190 it is needed to call :meth:`set_phase` to get the sampler of corresponding phase.
192 Args:
193 phases (str or list of str): All the phases of input.
194 datasets (Dataset or list of Dataset): All the dataset for each phase.
195 distribution (str, optional): Distribution of the negative items. Defaults to 'uniform'.
197 Attributes:
198 phase (str): the phase of sampler. It will not be set until :meth:`set_phase` is called.
199 """
201 def __init__(self, phases, datasets, distribution="uniform", alpha=1.0):
202 if not isinstance(phases, list):
203 phases = [phases]
204 if not isinstance(datasets, list):
205 datasets = [datasets]
206 if len(phases) != len(datasets):
207 raise ValueError(f"Phases {phases} and datasets {datasets} should have the same length.")
209 self.phases = phases
210 self.datasets = datasets
212 self.uid_field = datasets[0].uid_field
213 self.iid_field = datasets[0].iid_field
215 self.user_num = datasets[0].user_num
216 self.item_num = datasets[0].item_num
218 super().__init__(distribution=distribution, alpha=alpha)
220 def _get_candidates_list(self):
221 candidates_list = []
222 for dataset in self.datasets:
223 candidates_list.extend(dataset.inter_feat[self.iid_field].numpy())
224 return candidates_list
226 def _uni_sampling(self, sample_num):
227 return np.random.randint(1, self.item_num, sample_num)
229 def get_used_ids(self):
230 """Returns:
231 dict: Used item_ids is the same as positive item_ids.
232 Key is phase, and value is a numpy.ndarray which index is user_id, and element is a set of item_ids.
233 """
234 used_item_id = dict()
235 last = [set() for _ in range(self.user_num)]
236 for phase, dataset in zip(self.phases, self.datasets):
237 cur = np.array([set(s) for s in last])
238 for uid, iid in zip(
239 dataset.inter_feat[self.uid_field].numpy(),
240 dataset.inter_feat[self.iid_field].numpy(),
241 ):
242 cur[uid].add(iid)
243 last = used_item_id[phase] = cur
245 for used_item_set in used_item_id[self.phases[-1]]:
246 if len(used_item_set) + 1 == self.item_num: # [pad] is a item.
247 raise ValueError(
248 "Some users have interacted with all items, "
249 "which we can not sample negative items for them. "
250 "Please set `user_inter_num_interval` to filter those users."
251 )
252 return used_item_id
254 def set_phase(self, phase):
255 """Get the sampler of corresponding phase.
257 Args:
258 phase (str): The phase of new sampler.
260 Returns:
261 Sampler: the copy of this sampler, :attr:`phase` is set the same as input phase, and :attr:`used_ids`
262 is set to the value of corresponding phase.
263 """
264 if phase not in self.phases:
265 raise ValueError(f"Phase [{phase}] not exist.")
266 new_sampler = copy.copy(self)
267 new_sampler.phase = phase
268 new_sampler.used_ids = new_sampler.used_ids[phase]
269 return new_sampler
271 def sample_by_user_ids(self, user_ids, item_ids, num):
272 """Sampling by user_ids.
274 Args:
275 user_ids (numpy.ndarray or list): Input user_ids.
276 item_ids (numpy.ndarray or list): Input item_ids.
277 num (int): Number of sampled item_ids for each user_id.
279 Returns:
280 torch.tensor: Sampled item_ids.
281 item_ids[0], item_ids[len(user_ids)], item_ids[len(user_ids) * 2], ..., item_id[len(user_ids) * (num - 1)]
282 is sampled for user_ids[0];
283 item_ids[1], item_ids[len(user_ids) + 1], item_ids[len(user_ids) * 2 + 1], ...,
284 item_id[len(user_ids) * (num - 1) + 1] is sampled for user_ids[1]; ...; and so on.
285 """
286 try:
287 return self.sample_by_key_ids(user_ids, num)
288 except IndexError:
289 for user_id in user_ids:
290 if user_id < 0 or user_id >= self.user_num:
291 raise ValueError(f"user_id [{user_id}] not exist.")
294class KGSampler(AbstractSampler):
295 """:class:`KGSampler` is used to sample negative entities in a knowledge graph.
297 Args:
298 dataset (Dataset): The knowledge graph dataset, which contains triplets in a knowledge graph.
299 distribution (str, optional): Distribution of the negative entities. Defaults to 'uniform'.
300 """
302 def __init__(self, dataset, distribution="uniform", alpha=1.0):
303 self.dataset = dataset
305 self.hid_field = dataset.head_entity_field
306 self.tid_field = dataset.tail_entity_field
307 self.hid_list = dataset.head_entities
308 self.tid_list = dataset.tail_entities
310 self.head_entities = set(dataset.head_entities)
311 self.entity_num = dataset.entity_num
313 super().__init__(distribution=distribution, alpha=alpha)
315 def _uni_sampling(self, sample_num):
316 return np.random.randint(1, self.entity_num, sample_num)
318 def _get_candidates_list(self):
319 return list(self.hid_list) + list(self.tid_list)
321 def get_used_ids(self):
322 """Returns:
323 numpy.ndarray: Used entity_ids is the same as tail_entity_ids in knowledge graph.
324 Index is head_entity_id, and element is a set of tail_entity_ids.
325 """
326 used_tail_entity_id = np.array([set() for _ in range(self.entity_num)])
327 for hid, tid in zip(self.hid_list, self.tid_list):
328 used_tail_entity_id[hid].add(tid)
330 for used_tail_set in used_tail_entity_id:
331 if len(used_tail_set) + 1 == self.entity_num: # [pad] is a entity.
332 raise ValueError(
333 "Some head entities have relation with all entities, "
334 "which we can not sample negative entities for them."
335 )
336 return used_tail_entity_id
338 def sample_by_entity_ids(self, head_entity_ids, num=1):
339 """Sampling by head_entity_ids.
341 Args:
342 head_entity_ids (numpy.ndarray or list): Input head_entity_ids.
343 num (int, optional): Number of sampled entity_ids for each head_entity_id. Defaults to ``1``.
345 Returns:
346 torch.tensor: Sampled entity_ids.
347 entity_ids[0], entity_ids[len(head_entity_ids)], entity_ids[len(head_entity_ids) * 2], ...,
348 entity_id[len(head_entity_ids) * (num - 1)] is sampled for head_entity_ids[0];
349 entity_ids[1], entity_ids[len(head_entity_ids) + 1], entity_ids[len(head_entity_ids) * 2 + 1], ...,
350 entity_id[len(head_entity_ids) * (num - 1) + 1] is sampled for head_entity_ids[1]; ...; and so on.
351 """
352 try:
353 return self.sample_by_key_ids(head_entity_ids, num)
354 except IndexError:
355 for head_entity_id in head_entity_ids:
356 if head_entity_id not in self.head_entities:
357 raise ValueError(f"head_entity_id [{head_entity_id}] not exist.")
360class RepeatableSampler(AbstractSampler):
361 """:class:`RepeatableSampler` is used to sample negative items for each input user. The difference from
362 :class:`Sampler` is it can only sampling the items that have not appeared at all phases.
364 Args:
365 phases (str or list of str): All the phases of input.
366 dataset (Dataset): The union of all datasets for each phase.
367 distribution (str, optional): Distribution of the negative items. Defaults to 'uniform'.
369 Attributes:
370 phase (str): the phase of sampler. It will not be set until :meth:`set_phase` is called.
371 """
373 def __init__(self, phases, dataset, distribution="uniform", alpha=1.0):
374 if not isinstance(phases, list):
375 phases = [phases]
376 self.phases = phases
377 self.dataset = dataset
379 self.iid_field = dataset.iid_field
380 self.user_num = dataset.user_num
381 self.item_num = dataset.item_num
383 super().__init__(distribution=distribution, alpha=alpha)
385 def _uni_sampling(self, sample_num):
386 return np.random.randint(1, self.item_num, sample_num)
388 def _get_candidates_list(self):
389 return list(self.dataset.inter_feat[self.iid_field].numpy())
391 def get_used_ids(self):
392 """Returns:
393 numpy.ndarray: Used item_ids is the same as positive item_ids.
394 Index is user_id, and element is a set of item_ids.
395 """
396 return np.array([set() for _ in range(self.user_num)])
398 def sample_by_user_ids(self, user_ids, item_ids, num):
399 """Sampling by user_ids.
401 Args:
402 user_ids (numpy.ndarray or list): Input user_ids.
403 item_ids (numpy.ndarray or list): Input item_ids.
404 num (int): Number of sampled item_ids for each user_id.
406 Returns:
407 torch.tensor: Sampled item_ids.
408 item_ids[0], item_ids[len(user_ids)], item_ids[len(user_ids) * 2], ..., item_id[len(user_ids) * (num - 1)]
409 is sampled for user_ids[0];
410 item_ids[1], item_ids[len(user_ids) + 1], item_ids[len(user_ids) * 2 + 1], ...,
411 item_id[len(user_ids) * (num - 1) + 1] is sampled for user_ids[1]; ...; and so on.
412 """
413 try:
414 self.used_ids = np.array([{i} for i in item_ids])
415 return self.sample_by_key_ids(np.arange(len(user_ids)), num)
416 except IndexError:
417 for user_id in user_ids:
418 if user_id < 0 or user_id >= self.user_num:
419 raise ValueError(f"user_id [{user_id}] not exist.")
421 def set_phase(self, phase):
422 """Get the sampler of corresponding phase.
424 Args:
425 phase (str): The phase of new sampler.
427 Returns:
428 Sampler: the copy of this sampler, and :attr:`phase` is set the same as input phase.
429 """
430 if phase not in self.phases:
431 raise ValueError(f"Phase [{phase}] not exist.")
432 new_sampler = copy.copy(self)
433 new_sampler.phase = phase
434 return new_sampler
437class SeqSampler(AbstractSampler):
438 """:class:`SeqSampler` is used to sample negative item sequence.
440 Args:
441 datasets (Dataset or list of Dataset): All the dataset for each phase.
442 distribution (str, optional): Distribution of the negative items. Defaults to 'uniform'.
443 """
445 def __init__(self, dataset, distribution="uniform", alpha=1.0):
446 self.dataset = dataset
448 self.iid_field = dataset.iid_field
449 self.user_num = dataset.user_num
450 self.item_num = dataset.item_num
452 super().__init__(distribution=distribution, alpha=alpha)
454 def _uni_sampling(self, sample_num):
455 return np.random.randint(1, self.item_num, sample_num)
457 def get_used_ids(self):
458 pass
460 def sample_neg_sequence(self, pos_sequence):
461 """For each moment, sampling one item from all the items except the one the user clicked on at that moment.
463 Args:
464 pos_sequence (torch.Tensor): all users' item history sequence, with the shape of `(N, )`.
466 Returns:
467 torch.tensor : all users' negative item history sequence.
469 """
470 total_num = len(pos_sequence)
471 value_ids = np.zeros(total_num, dtype=np.int64)
472 check_list = np.arange(total_num)
473 while len(check_list) > 0:
474 value_ids[check_list] = self.sampling(len(check_list))
475 check_index = np.where(value_ids[check_list] == pos_sequence[check_list])
476 check_list = check_list[check_index]
478 return torch.tensor(value_ids)