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

1# @Author : Yupeng Hou 

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

3# @File : sampler.py 

4 

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 

9 

10"""hopwise.sampler 

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

12""" 

13 

14import copy 

15from collections import Counter 

16 

17import numpy as np 

18import torch 

19 

20 

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. 

25 

26 Args: 

27 distribution (str): The string of distribution, which is used for subclass. 

28 

29 Attributes: 

30 used_ids (numpy.ndarray): The result of :meth:`get_used_ids`. 

31 """ 

32 

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

38 

39 def set_distribution(self, distribution): 

40 """Set the distribution of sampler. 

41 

42 Args: 

43 distribution (str): Distribution of the negative items. 

44 """ 

45 self.distribution = distribution 

46 if distribution == "popularity": 

47 self._build_alias_table() 

48 

49 def _uni_sampling(self, sample_num): 

50 """Sample [sample_num] items in the uniform distribution. 

51 

52 Args: 

53 sample_num (int): the number of samples. 

54 

55 Returns: 

56 sample_list (np.array): a list of samples. 

57 """ 

58 raise NotImplementedError("Method [_uni_sampling] should be implemented") 

59 

60 def _get_candidates_list(self): 

61 """Get sample candidates list for _pop_sampling() 

62 

63 Returns: 

64 candidates_list (list): a list of candidates id. 

65 """ 

66 raise NotImplementedError("Method [_get_candidates_list] should be implemented") 

67 

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) 

95 

96 def _pop_sampling(self, sample_num): 

97 """Sample [sample_num] items in the popularity-biased distribution. 

98 

99 Args: 

100 sample_num (int): the number of samples. 

101 

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 = [] 

109 

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

115 

116 return np.array(final_random_list) 

117 

118 def sampling(self, sample_num): 

119 """Sampling [sample_num] item_ids. 

120 

121 Args: 

122 sample_num (int): the number of samples. 

123 

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

133 

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

139 

140 def sample_by_key_ids(self, key_ids, num): 

141 """Sampling by key_ids. 

142 

143 Args: 

144 key_ids (numpy.ndarray or list): Input key_ids. 

145 num (int): Number of sampled value_ids for each key_id. 

146 

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) 

184 

185 

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. 

191 

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

196 

197 Attributes: 

198 phase (str): the phase of sampler. It will not be set until :meth:`set_phase` is called. 

199 """ 

200 

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

208 

209 self.phases = phases 

210 self.datasets = datasets 

211 

212 self.uid_field = datasets[0].uid_field 

213 self.iid_field = datasets[0].iid_field 

214 

215 self.user_num = datasets[0].user_num 

216 self.item_num = datasets[0].item_num 

217 

218 super().__init__(distribution=distribution, alpha=alpha) 

219 

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 

225 

226 def _uni_sampling(self, sample_num): 

227 return np.random.randint(1, self.item_num, sample_num) 

228 

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 

244 

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 

253 

254 def set_phase(self, phase): 

255 """Get the sampler of corresponding phase. 

256 

257 Args: 

258 phase (str): The phase of new sampler. 

259 

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 

270 

271 def sample_by_user_ids(self, user_ids, item_ids, num): 

272 """Sampling by user_ids. 

273 

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. 

278 

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

292 

293 

294class KGSampler(AbstractSampler): 

295 """:class:`KGSampler` is used to sample negative entities in a knowledge graph. 

296 

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

301 

302 def __init__(self, dataset, distribution="uniform", alpha=1.0): 

303 self.dataset = dataset 

304 

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 

309 

310 self.head_entities = set(dataset.head_entities) 

311 self.entity_num = dataset.entity_num 

312 

313 super().__init__(distribution=distribution, alpha=alpha) 

314 

315 def _uni_sampling(self, sample_num): 

316 return np.random.randint(1, self.entity_num, sample_num) 

317 

318 def _get_candidates_list(self): 

319 return list(self.hid_list) + list(self.tid_list) 

320 

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) 

329 

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 

337 

338 def sample_by_entity_ids(self, head_entity_ids, num=1): 

339 """Sampling by head_entity_ids. 

340 

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

344 

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

358 

359 

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. 

363 

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

368 

369 Attributes: 

370 phase (str): the phase of sampler. It will not be set until :meth:`set_phase` is called. 

371 """ 

372 

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 

378 

379 self.iid_field = dataset.iid_field 

380 self.user_num = dataset.user_num 

381 self.item_num = dataset.item_num 

382 

383 super().__init__(distribution=distribution, alpha=alpha) 

384 

385 def _uni_sampling(self, sample_num): 

386 return np.random.randint(1, self.item_num, sample_num) 

387 

388 def _get_candidates_list(self): 

389 return list(self.dataset.inter_feat[self.iid_field].numpy()) 

390 

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

397 

398 def sample_by_user_ids(self, user_ids, item_ids, num): 

399 """Sampling by user_ids. 

400 

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. 

405 

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

420 

421 def set_phase(self, phase): 

422 """Get the sampler of corresponding phase. 

423 

424 Args: 

425 phase (str): The phase of new sampler. 

426 

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 

435 

436 

437class SeqSampler(AbstractSampler): 

438 """:class:`SeqSampler` is used to sample negative item sequence. 

439 

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

444 

445 def __init__(self, dataset, distribution="uniform", alpha=1.0): 

446 self.dataset = dataset 

447 

448 self.iid_field = dataset.iid_field 

449 self.user_num = dataset.user_num 

450 self.item_num = dataset.item_num 

451 

452 super().__init__(distribution=distribution, alpha=alpha) 

453 

454 def _uni_sampling(self, sample_num): 

455 return np.random.randint(1, self.item_num, sample_num) 

456 

457 def get_used_ids(self): 

458 pass 

459 

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. 

462 

463 Args: 

464 pos_sequence (torch.Tensor): all users' item history sequence, with the shape of `(N, )`. 

465 

466 Returns: 

467 torch.tensor : all users' negative item history sequence. 

468 

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] 

477 

478 return torch.tensor(value_ids)