Coverage for hopwise/data/dataset/customized_dataset.py: 55%

393 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1# @Time : 2020/10/19 

2# @Author : Yupeng Hou 

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

4 

5# UPDATE 

6# @Time : 2021/7/9 

7# @Author : Yupeng Hou 

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

9 

10"""hopwise.data.customized_dataset 

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

12 

13We only recommend building customized datasets by inheriting. 

14 

15Customized datasets named ``[Model Name]Dataset`` can be automatically called. 

16""" 

17 

18import datetime 

19 

20import numba 

21import numpy as np 

22import pandas as pd 

23import torch 

24from sklearn.mixture import GaussianMixture as GMM 

25 

26from hopwise.data.dataset import ( 

27 KGSeqDataset, 

28 KnowledgeBasedDataset, 

29 KnowledgePathDataset, 

30 SequentialDataset, 

31 UserItemKnowledgePathDataset, 

32) 

33from hopwise.data.dataset.kg_path_dataset import CSRGraph 

34from hopwise.data.interaction import Interaction 

35from hopwise.sampler import SeqSampler 

36from hopwise.utils import FeatureType, progress_bar, set_color 

37 

38 

39class GRU4RecKGDataset(KGSeqDataset): 

40 def __init__(self, config): 

41 super().__init__(config) 

42 

43 

44class KSRDataset(KGSeqDataset): 

45 def __init__(self, config): 

46 super().__init__(config) 

47 

48 

49class DIENDataset(SequentialDataset): 

50 """:class:`DIENDataset` is based on :class:`~hopwise.data.dataset.sequential_dataset.SequentialDataset`. 

51 It is different from :class:`SequentialDataset` in `data_augmentation`. 

52 It add users' negative item list to interaction. 

53 

54 The original version of sampling negative item list is implemented by Zhichao Feng (fzcbupt@gmail.com) in 

55 2021/2/25, and he updated the codes in 2021/3/19. In 2021/7/9, 

56 Yupeng refactored SequentialDataset & SequentialDataLoader, then refactored DIENDataset, either. 

57 

58 Attributes: 

59 augmentation (bool): Whether the interactions should be augmented in hopwise. 

60 seq_sample (hopwise.sampler.SeqSampler): A sampler used to sample negative item sequence. 

61 neg_item_list_field (str): Field name for negative item sequence. 

62 neg_item_list (torch.tensor): all users' negative item history sequence. 

63 """ 

64 

65 def __init__(self, config): 

66 super().__init__(config) 

67 

68 list_suffix = config["LIST_SUFFIX"] 

69 neg_prefix = config["NEG_PREFIX"] 

70 self.seq_sampler = SeqSampler(self) 

71 self.neg_item_list_field = neg_prefix + self.iid_field + list_suffix 

72 self.neg_item_list = self.seq_sampler.sample_neg_sequence(self.inter_feat[self.iid_field]) 

73 

74 def data_augmentation(self): 

75 """Augmentation processing for sequential dataset. 

76 

77 E.g., ``u1`` has purchase sequence ``<i1, i2, i3, i4>``, 

78 then after augmentation, we will generate three cases. 

79 

80 ``u1, <i1> | i2`` 

81 

82 (Which means given user_id ``u1`` and item_seq ``<i1>``, 

83 we need to predict the next item ``i2``.) 

84 

85 The other cases are below: 

86 

87 ``u1, <i1, i2> | i3`` 

88 

89 ``u1, <i1, i2, i3> | i4`` 

90 """ 

91 self.logger.debug("data_augmentation") 

92 

93 self._aug_presets() 

94 

95 self._check_field("uid_field", "time_field") 

96 max_item_list_len = self.config["MAX_ITEM_LIST_LENGTH"] 

97 self.sort(by=[self.uid_field, self.time_field], ascending=True) 

98 last_uid = None 

99 uid_list, item_list_index, target_index, item_list_length = [], [], [], [] 

100 seq_start = 0 

101 for i, uid in enumerate(self.inter_feat[self.uid_field].numpy()): 

102 if last_uid != uid: 

103 last_uid = uid 

104 seq_start = i 

105 else: 

106 if i - seq_start > max_item_list_len: 

107 seq_start += 1 

108 uid_list.append(uid) 

109 item_list_index.append(slice(seq_start, i)) 

110 target_index.append(i) 

111 item_list_length.append(i - seq_start) 

112 

113 uid_list = np.array(uid_list) 

114 item_list_index = np.array(item_list_index) 

115 target_index = np.array(target_index) 

116 item_list_length = np.array(item_list_length, dtype=np.int64) 

117 

118 new_length = len(item_list_index) 

119 new_data = self.inter_feat[target_index] 

120 new_dict = { 

121 self.item_list_length_field: torch.tensor(item_list_length), 

122 } 

123 

124 for field in self.inter_feat: 

125 if field != self.uid_field: 

126 list_field = getattr(self, f"{field}_list_field") 

127 list_len = self.field2seqlen[list_field] 

128 shape = (new_length, list_len) if isinstance(list_len, int) else (new_length,) + list_len 

129 if ( 

130 self.field2type[field] in [FeatureType.FLOAT, FeatureType.FLOAT_SEQ] 

131 and field in self.config["numerical_features"] 

132 ): 

133 shape += (2,) 

134 list_ftype = self.field2type[list_field] 

135 dtype = torch.int64 if list_ftype in [FeatureType.TOKEN, FeatureType.TOKEN_SEQ] else torch.float64 

136 new_dict[list_field] = torch.zeros(shape, dtype=dtype) 

137 

138 value = self.inter_feat[field] 

139 for i, (index, length) in enumerate(zip(item_list_index, item_list_length)): 

140 new_dict[list_field][i][:length] = value[index] 

141 

142 # DIEN 

143 if field == self.iid_field: 

144 new_dict[self.neg_item_list_field] = torch.zeros(shape, dtype=dtype) 

145 for i, (index, length) in enumerate(zip(item_list_index, item_list_length)): 

146 new_dict[self.neg_item_list_field][i][:length] = self.neg_item_list[index] 

147 

148 new_data.update(Interaction(new_dict)) 

149 self.inter_feat = new_data 

150 

151 

152class KGGLMDatasetMixin: 

153 """Mixin class containing KGGLM-specific dataset logic. 

154 

155 This mixin should be used with KnowledgePathDataset or UserItemKnowledgePathDataset 

156 to create the appropriate KGGLM dataset class. 

157 """ 

158 

159 def _get_field_from_config(self): 

160 super()._get_field_from_config() 

161 self.train_stage = self.config["train_stage"] 

162 

163 path_sample_args = self.config["path_sample_args"] 

164 self.pretrain_hop_length = path_sample_args["pretrain_hop_length"] 

165 if isinstance(self.pretrain_hop_length, str): 

166 self.pretrain_hop_length = tuple(map(int, self.pretrain_hop_length[1:-1].split(","))) 

167 self.pretrain_paths = path_sample_args["pretrain_paths"] 

168 

169 def generate_user_path_dataset(self): 

170 if self.train_stage == "pretrain": 

171 self.generate_pretrain_dataset() 

172 else: 

173 super().generate_user_path_dataset() 

174 

175 def generate_pretrain_dataset(self): 

176 """Generate pretrain dataset for KGGLM model using CSR-based parallel random walks.""" 

177 

178 if self._path_dataset is None: 

179 csr_matrix = self._create_ckg_sparse_matrix(form="csr", show_relation=True) 

180 csr_graph = CSRGraph.from_sparse_matrix(csr_matrix) 

181 indptr, indices, relations = csr_graph.unpack() 

182 

183 # UI relations excluded (weight=0) 

184 ui_rel_id = self.relation_num - 1 

185 weights = np.where(relations == ui_rel_id, 0.0, 1.0).astype(np.float32) 

186 

187 min_hop, max_hop = self.pretrain_hop_length 

188 max_tries_per_entity = self.config["path_sample_args"]["MAX_RW_TRIES_PER_IID"] 

189 

190 entity_ids = self._get_entity_ids_range() 

191 

192 # Generate start nodes: each entity gets pretrain_paths * max_tries samples 

193 samples_per_entity = self.pretrain_paths * max_tries_per_entity 

194 all_start_nodes = np.repeat(entity_ids, samples_per_entity) 

195 

196 # Generate random hop lengths for each walk 

197 all_hop_lengths = np.random.randint(min_hop, max_hop + 1, size=len(all_start_nodes)) 

198 

199 self.logger.info( 

200 set_color(f"Running {len(all_start_nodes)} parallel random walks for pretraining...", "blue") 

201 ) 

202 

203 # Run parallel random walks with max_hop (we'll truncate based on actual hop length later) 

204 paths, path_rels = _kgglm_csr_parallel_random_walks( 

205 indptr, indices, relations, weights, all_start_nodes, all_hop_lengths, max_hop 

206 ) 

207 

208 # Deduplicate and collect unique paths per entity 

209 unique_paths = set() 

210 entity_path_counts = {} 

211 

212 iter_paths = progress_bar( 

213 range(len(paths)), 

214 ncols=100, 

215 total=len(paths), 

216 desc=set_color("KGGLM Pre-training Path Sampling", "red", progress=True), 

217 ) 

218 

219 for idx in iter_paths: 

220 entity = all_start_nodes[idx] 

221 hop_length = all_hop_lengths[idx] 

222 

223 # Check if entity already has enough paths 

224 if entity_path_counts.get(entity, 0) >= self.pretrain_paths: 

225 continue 

226 

227 # Get the actual path (truncated to hop_length) 

228 path = tuple(paths[idx, : hop_length + 1]) 

229 path_rel = tuple(path_rels[idx, :hop_length]) 

230 

231 # Skip if path has invalid nodes (walk got stuck) 

232 if -1 in path: 

233 continue 

234 

235 # Build path with relations interleaved 

236 path_with_rel = [] 

237 for i, node in enumerate(path): 

238 path_with_rel.append(node) 

239 if i < len(path_rel): 

240 path_with_rel.append(path_rel[i]) 

241 path_with_rel = tuple(path_with_rel) 

242 

243 # Add to unique paths if not seen 

244 if path_with_rel not in unique_paths: 

245 unique_paths.add(path_with_rel) 

246 entity_path_counts[entity] = entity_path_counts.get(entity, 0) + 1 

247 

248 # Format paths to string 

249 unique_paths = sorted(list(unique_paths)) # Sort for consistency 

250 formatted_paths = [self._format_path(np.array(path)) for path in unique_paths] 

251 path_string = "\n".join(formatted_paths) 

252 

253 self._path_dataset = path_string 

254 

255 

256@numba.njit(parallel=True) 

257def _kgglm_csr_parallel_random_walks(indptr, indices, relations, weights, start_nodes, hop_lengths, max_hop): 

258 """Parallel random walks on CSR graph with variable hop lengths for KGGLM pretraining. 

259 

260 Note: Set np.random.seed() before calling this function for reproducibility. 

261 

262 Args: 

263 indptr: CSR row pointers 

264 indices: CSR column indices 

265 relations: Edge relation types 

266 weights: Edge weights 

267 start_nodes: Array of starting nodes 

268 hop_lengths: Array of hop lengths per walk 

269 max_hop: Maximum hop length (for output array sizing) 

270 

271 Returns: 

272 paths: (n_walks, max_hop + 1) node paths 

273 path_relations: (n_walks, max_hop) relation paths 

274 """ 

275 n_walks = len(start_nodes) 

276 paths = np.full((n_walks, max_hop + 1), -1, dtype=np.int64) 

277 path_relations = np.full((n_walks, max_hop), -1, dtype=np.int64) 

278 

279 for i in numba.prange(n_walks): 

280 node = start_nodes[i] 

281 num_steps = hop_lengths[i] 

282 paths[i, 0] = node 

283 

284 for step in range(num_steps): 

285 start_idx = indptr[node] 

286 end_idx = indptr[node + 1] 

287 n_neighbors = end_idx - start_idx 

288 

289 if n_neighbors == 0: 

290 break 

291 

292 # random choice replacement for numba with p=weights 

293 edge_weights = weights[start_idx:end_idx] 

294 total_weight = edge_weights.sum() 

295 

296 if total_weight == 0: 

297 break 

298 

299 r = np.random.random() * total_weight 

300 cumsum = 0.0 

301 selected_idx = 0 

302 for j in range(n_neighbors): 

303 cumsum += edge_weights[j] 

304 if r <= cumsum: 

305 selected_idx = j 

306 break 

307 

308 neighbor_idx = start_idx + selected_idx 

309 node = indices[neighbor_idx] 

310 paths[i, step + 1] = node 

311 path_relations[i, step] = relations[neighbor_idx] 

312 

313 return paths, path_relations 

314 

315 

316class KGGLMDataset(KGGLMDatasetMixin, KnowledgePathDataset): 

317 """KGGLM dataset inheriting from KnowledgePathDataset.""" 

318 

319 def _get_entity_ids_range(self): 

320 """Get the range of entity IDs in the graph, which is used as the starting point for random walks. 

321 In this case, we only consider item entities, so the minimum ID is 1 + number of users. 

322 """ 

323 graph_min_iid = 1 + self.user_num 

324 num_entities = self.entity_num 

325 return np.arange(graph_min_iid, self.user_num + num_entities, dtype=np.int64) 

326 

327 

328class UserItemKGGLMDataset(KGGLMDatasetMixin, UserItemKnowledgePathDataset): 

329 """KGGLM dataset inheriting from UserItemKnowledgePathDataset. 

330 

331 Used when both user and item knowledge graph links are available. 

332 """ 

333 

334 def _get_entity_ids_range(self): 

335 """Get the range of entity IDs in the graph, which is used as the starting point for random walks. 

336 In this case, we consider both user and item entities, so the minimum ID is 1 (skipping ID 0) and ignore 

337 the padding item with ID 0 as well. 

338 """ 

339 graph_min_iid = 1 

340 item_min_iid = 1 + self.user_num 

341 num_entities = self.entity_num 

342 entity_ids = np.concatenate( 

343 [ 

344 np.arange(graph_min_iid, self.user_num, dtype=np.int64), 

345 np.arange(item_min_iid, num_entities, dtype=np.int64), 

346 ] 

347 ) 

348 return entity_ids 

349 

350 

351class TPRecTimestampDataset: 

352 """ 

353 A class to create a clustered dataset based on interaction timestamp specifically for the TPRec model 

354 """ 

355 

356 Y = 2000 

357 seasons = [ 

358 (0, (datetime.date(Y, 1, 1), datetime.date(Y, 3, 20))), # 'winter' 

359 (1, (datetime.date(Y, 3, 21), datetime.date(Y, 6, 20))), # 'spring' 

360 (2, (datetime.date(Y, 6, 21), datetime.date(Y, 9, 22))), # 'summer' 

361 (3, (datetime.date(Y, 9, 23), datetime.date(Y, 12, 20))), # 'autumn' 

362 (0, (datetime.date(Y, 12, 21), datetime.date(Y, 12, 31))), 

363 ] # 'winter' 

364 

365 def __init__(self, config, inter_feat, set="train", gmm=None): 

366 self.config = config 

367 self.inter_feat = inter_feat 

368 self.set = set 

369 data = {"users": inter_feat.user_id, "item": inter_feat.item_id, "timestamps": inter_feat.timestamp} 

370 self.data = pd.DataFrame(data) 

371 self.data.timestamps = self.data.timestamps.astype(int) 

372 self.data.timestamps = pd.to_datetime(self.data.timestamps, unit="s").dt.date 

373 

374 self.user_item_timestamp = np.array(self.data.timestamps) 

375 time2num = self._timeanalysis() 

376 

377 if config["cluster_feature"] == "all": 

378 fileTime = self._get_all_cluster_feature(time2num) 

379 elif config["cluster_feature"] == "w-stat": 

380 fileTime = self._get_w_stat_cluster_feature() 

381 elif config["cluster_feature"] == "w-stru": 

382 fileTime = self._get_w_stru_cluster_feature(time2num) 

383 else: 

384 raise ValueError( 

385 f"Unsupported cluster_feature: {config['cluster_feature']}. " 

386 "Available options are 'all', 'w-stat', 'w-stru'." 

387 ) 

388 

389 if self.set == "train": 

390 gmmModel, timeNum, labels, timeClassifyLabel = self._hierarchicalTime(fileTime) 

391 self.gmm_model = gmmModel 

392 self.timenum = timeNum 

393 self.timeClassifyLabel = timeClassifyLabel 

394 else: 

395 labels = self.test_knn_cluster(gmm, config["cluster_feature"]) 

396 

397 ucp_hash = self._generate_clus_dict(labels) 

398 uc_weight = self._generate_user_agent_num(ucp_hash) 

399 

400 self.uc_weight = uc_weight 

401 

402 def test_knn_cluster(self, gmm, cluster_feature): 

403 if cluster_feature == "all": 

404 fileTime = pd.DataFrame( 

405 self.data, 

406 columns=[ 

407 "pur_frequancy", 

408 "order1_90", 

409 "order2_90", 

410 "order1_30", 

411 "order2_30", 

412 "order1_7", 

413 "order2_7", 

414 "order1_1", 

415 "order2_1", 

416 "tfa_year", 

417 "tfa_month", 

418 "tfa_day", 

419 "tfa_weekday", 

420 "tfa_weekday_1", 

421 "tfa_weekday_2", 

422 "tfa_weekday_3", 

423 "tfa_weekday_4", 

424 "tfa_weekday_5", 

425 "tfa_weekday_6", 

426 "tfa_weekday_7", 

427 "tfa_season", 

428 "tfa_season_0", 

429 "tfa_season_1", 

430 "tfa_season_2", 

431 "tfa_season_3", 

432 ], 

433 ) 

434 fileTime["tfa_year"] = fileTime["tfa_year"] - fileTime["tfa_year"].min() 

435 elif cluster_feature == "w-stat": 

436 fileTime = pd.DataFrame( 

437 self.data, 

438 columns=[ 

439 "tfa_year", 

440 "tfa_month", 

441 "tfa_day", 

442 "tfa_weekday", 

443 "tfa_weekday_1", 

444 "tfa_weekday_2", 

445 "tfa_weekday_3", 

446 "tfa_weekday_4", 

447 "tfa_weekday_5", 

448 "tfa_weekday_6", 

449 "tfa_weekday_7", 

450 "tfa_season", 

451 "tfa_season_0", 

452 "tfa_season_1", 

453 "tfa_season_2", 

454 "tfa_season_3", 

455 ], 

456 ) 

457 fileTime["tfa_year"] = fileTime["tfa_year"] - fileTime["tfa_year"].min() 

458 elif cluster_feature == "w-stru": 

459 fileTime = pd.DataFrame( 

460 self.data, 

461 columns=[ 

462 "pur_frequancy", 

463 "order1_90", 

464 "order2_90", 

465 "order1_30", 

466 "order2_30", 

467 "order1_7", 

468 "order2_7", 

469 "order1_1", 

470 "order2_1", 

471 ], 

472 ) 

473 

474 x = np.array(fileTime) 

475 x = (x - x.min(axis=0)) / (x.max(axis=0) - x.min(axis=0)) 

476 test_cluster_label = gmm.predict(x) 

477 

478 return test_cluster_label 

479 

480 def _get_all_cluster_feature(self, time2num): 

481 user_item_timestamp = np.array(self.user_item_timestamp) 

482 # =============================== Structural Features ================================= [90, 30, 7, 1] 

483 add_df = pd.DataFrame( 

484 columns=[ 

485 "pur_frequancy", 

486 "order1_90", 

487 "order2_90", 

488 "order1_30", 

489 "order2_30", 

490 "order1_7", 

491 "order2_7", 

492 "order1_1", 

493 "order2_1", 

494 ], 

495 data=np.array([time2num[i] for i in user_item_timestamp]), 

496 ) 

497 ( 

498 self.data["pur_frequancy"], 

499 self.data["order1_90"], 

500 self.data["order2_90"], 

501 self.data["order1_30"], 

502 self.data["order2_30"], 

503 self.data["order1_7"], 

504 self.data["order2_7"], 

505 self.data["order1_1"], 

506 self.data["order2_1"], 

507 ) = ( 

508 add_df["pur_frequancy"], 

509 add_df["order1_90"], 

510 add_df["order2_90"], 

511 add_df["order1_30"], 

512 add_df["order2_30"], 

513 add_df["order1_7"], 

514 add_df["order2_7"], 

515 add_df["order1_1"], 

516 add_df["order2_1"], 

517 ) 

518 

519 # =============================== Stastical Features ================================= 

520 self.data["tfa_year"] = np.array([x.year for x in self.data.timestamps]) 

521 self.data["tfa_month"] = np.array([x.month for x in self.data.timestamps]) 

522 self.data["tfa_day"] = np.array([x.day for x in self.data.timestamps]) 

523 self.data["tfa_weekday"] = np.array([x.isoweekday() for x in self.data.timestamps]) 

524 

525 tfa_weekday = pd.get_dummies(self.data.tfa_weekday, prefix="tfa_weekday") # one hot encoding 

526 self.data = pd.concat((self.data, tfa_weekday), axis=1) 

527 

528 self.data["tfa_season"] = np.array([self._get_season(x) for x in self.data.timestamps]) 

529 tfa_season = pd.get_dummies(self.data.tfa_season, prefix="tfa_season") # one hot encoding 

530 self.data = pd.concat((self.data, tfa_season), axis=1) 

531 

532 fileTime = pd.DataFrame( 

533 self.data, 

534 columns=[ 

535 "pur_frequancy", 

536 "order1_90", 

537 "order2_90", 

538 "order1_30", 

539 "order2_30", 

540 "order1_7", 

541 "order2_7", 

542 "order1_1", 

543 "order2_1", 

544 "tfa_year", 

545 "tfa_month", 

546 "tfa_day", 

547 "tfa_weekday", 

548 "tfa_weekday_1", 

549 "tfa_weekday_2", 

550 "tfa_weekday_3", 

551 "tfa_weekday_4", 

552 "tfa_weekday_5", 

553 "tfa_weekday_6", 

554 "tfa_weekday_7", 

555 "tfa_season", 

556 "tfa_season_0", 

557 "tfa_season_1", 

558 "tfa_season_2", 

559 "tfa_season_3", 

560 ], 

561 ) 

562 fileTime["tfa_year"] = fileTime["tfa_year"] - fileTime["tfa_year"].min() 

563 

564 return fileTime 

565 

566 def _get_w_stat_cluster_feature(self): 

567 self.data["tfa_year"] = np.array([x.year for x in self.data.timestamps]) 

568 self.data["tfa_month"] = np.array([x.month for x in self.data.timestamps]) 

569 self.data["tfa_day"] = np.array([x.day for x in self.data.timestamps]) 

570 self.data["tfa_weekday"] = np.array([x.isoweekday() for x in self.data.timestamps]) 

571 

572 tfa_weekday = pd.get_dummies(self.data.tfa_weekday, prefix="tfa_weekday") # one hot encoding 

573 self.data = pd.concat((self.data, tfa_weekday), axis=1) 

574 

575 self.data["tfa_season"] = np.array([self._get_season(x) for x in self.data.timestamps]) 

576 tfa_season = pd.get_dummies(self.data.tfa_season, prefix="tfa_season") # one hot encoding 

577 self.data = pd.concat((self.data, tfa_season), axis=1) 

578 

579 fileTime = pd.DataFrame( 

580 self.data, 

581 columns=[ 

582 "tfa_year", 

583 "tfa_month", 

584 "tfa_day", 

585 "tfa_weekday", 

586 "tfa_weekday_1", 

587 "tfa_weekday_2", 

588 "tfa_weekday_3", 

589 "tfa_weekday_4", 

590 "tfa_weekday_5", 

591 "tfa_weekday_6", 

592 "tfa_weekday_7", 

593 "tfa_season", 

594 "tfa_season_0", 

595 "tfa_season_1", 

596 "tfa_season_2", 

597 "tfa_season_3", 

598 ], 

599 ) 

600 fileTime["tfa_year"] = fileTime["tfa_year"] - fileTime["tfa_year"].min() 

601 

602 return fileTime 

603 

604 def _get_w_stru_cluster_feature(self, time2num): 

605 add_df = pd.DataFrame( 

606 columns=[ 

607 "pur_frequancy", 

608 "order1_90", 

609 "order2_90", 

610 "order1_30", 

611 "order2_30", 

612 "order1_7", 

613 "order2_7", 

614 "order1_1", 

615 "order2_1", 

616 ], 

617 data=np.array([time2num[i] for i in self.user_item_timestamp]), 

618 ) 

619 ( 

620 self.data["pur_frequancy"], 

621 self.data["order1_90"], 

622 self.data["order2_90"], 

623 self.data["order1_30"], 

624 self.data["order2_30"], 

625 self.data["order1_7"], 

626 self.data["order2_7"], 

627 self.data["order1_1"], 

628 self.data["order2_1"], 

629 ) = ( 

630 add_df["pur_frequancy"], 

631 add_df["order1_90"], 

632 add_df["order2_90"], 

633 add_df["order1_30"], 

634 add_df["order2_30"], 

635 add_df["order1_7"], 

636 add_df["order2_7"], 

637 add_df["order1_1"], 

638 add_df["order2_1"], 

639 ) 

640 

641 fileTime = pd.DataFrame( 

642 self.data, 

643 columns=[ 

644 "pur_frequancy", 

645 "order1_90", 

646 "order2_90", 

647 "order1_30", 

648 "order2_30", 

649 "order1_7", 

650 "order2_7", 

651 "order1_1", 

652 "order2_1", 

653 ], 

654 ) 

655 

656 return fileTime 

657 

658 def _timeanalysis(self): 

659 dac_time = self.data.timestamps.value_counts() 

660 dac_time_date = pd.to_datetime(dac_time.index) 

661 

662 dac_time_day = dac_time_date - dac_time_date.min() 

663 time2num = {} 

664 time2relative = {} 

665 serial2PrefixSum = {} 

666 for i in range(len(dac_time)): 

667 # quante volte occorrono i timestamp? 

668 time2num[dac_time.index[i]] = [dac_time.values[i]] 

669 for i in range(len(dac_time)): 

670 time2relative[dac_time_day.days[i]] = dac_time.index[i] 

671 mapIndex = sorted(time2relative.keys()) 

672 serial2PrefixSum[0] = 0 

673 for i in range(1, mapIndex[-1] + 1): 

674 cur = time2num.get(time2relative.get(i, 0), 0) 

675 if cur: 

676 serial2PrefixSum[i] = serial2PrefixSum[i - 1] + cur[0] 

677 else: 

678 serial2PrefixSum[i] = serial2PrefixSum[i - 1] 

679 

680 for gap in [90, 30, 7, 1]: 

681 self._structuralWithGap(gap, time2num, time2relative, mapIndex, serial2PrefixSum) 

682 return time2num 

683 

684 def _hierarchicalTime(self, filetime): 

685 ui2label = {} 

686 x = np.array(filetime) 

687 x = (x - x.min(axis=0)) / (x.max(axis=0) - x.min(axis=0)) 

688 

689 models = GMM(self.config["cluster_num"], covariance_type="full", random_state=self.config["seed"]).fit(x) 

690 labels = models.predict(x) 

691 

692 for user, item, label in zip(self.data.users, self.data.item, labels): 

693 ui2label[(user, item)] = label 

694 

695 return models, self.config["cluster_num"], labels, ui2label 

696 

697 def _generate_clus_dict(self, clus_label): 

698 uid_pid_clu = pd.DataFrame(self.data, columns=["users", "item"]) 

699 uid_pid_clu["clu_label"] = clus_label 

700 

701 uid_pid_clu_list = uid_pid_clu.values.tolist() 

702 ucp_hash = {} 

703 # depending on the server load this can take a long time 

704 for [uid, pid, clu] in progress_bar(uid_pid_clu_list, desc=f"generating clusters dict {self.set}"): 

705 if uid not in ucp_hash: 

706 # ucp_hash : {uids{c1: pid, c2:pid, ...}, ...} 

707 ucp_hash[uid] = {clu: [pid]} 

708 else: 

709 if clu not in ucp_hash[uid]: 

710 ucp_hash[uid][clu] = [] 

711 ucp_hash[uid][clu].append(pid) 

712 

713 return ucp_hash 

714 

715 def _generate_user_agent_num(self, ucp_hash): 

716 u_c_weight = ucp_hash 

717 # depending on the server load this can take a long time 

718 for u in progress_bar(u_c_weight, desc=f"generating user agent number {self.set}"): 

719 tmp_u_tot = 0 

720 for c in u_c_weight[u]: 

721 tmp_u_tot = tmp_u_tot + len(u_c_weight[u][c]) 

722 for c in u_c_weight[u]: 

723 u_c_weight[u][c] = len(u_c_weight[u][c]) / tmp_u_tot 

724 return u_c_weight 

725 

726 def _structuralWithGap(self, gap, time2Num, time2relative, mapIndex, serial2PrefixSum): 

727 # don't ask what this function does. I don't know. 

728 second_order_serial2PrefixSum = {} 

729 init_left = (serial2PrefixSum[2 * gap] - serial2PrefixSum[0] - 2 * serial2PrefixSum[gap]) / gap 

730 for i in range(mapIndex[-1] + 1): 

731 if i <= 2 * gap: 

732 second_order_serial2PrefixSum[i] = init_left 

733 else: 

734 second_order_serial2PrefixSum[i] = ( 

735 serial2PrefixSum[i] - 2 * serial2PrefixSum[i - gap] + serial2PrefixSum[i - 2 * gap] 

736 ) / gap 

737 

738 init_left_2 = ( 

739 second_order_serial2PrefixSum[2 * gap] 

740 - second_order_serial2PrefixSum[0] 

741 - 2 * second_order_serial2PrefixSum[gap] 

742 ) / gap 

743 for idx in mapIndex: 

744 if idx <= 2 * gap: 

745 gap_left = init_left 

746 gap_left_2 = init_left_2 

747 else: 

748 gap_left = ( 

749 serial2PrefixSum[idx] - 2 * serial2PrefixSum[idx - gap] + serial2PrefixSum[idx - 2 * gap] 

750 ) / gap 

751 gap_left_2 = ( 

752 second_order_serial2PrefixSum[idx] 

753 - 2 * second_order_serial2PrefixSum[idx - gap] 

754 + second_order_serial2PrefixSum[idx - 2 * gap] 

755 ) / gap 

756 

757 time2Num[time2relative[idx]].append(gap_left) 

758 time2Num[time2relative[idx]].append(gap_left_2) 

759 

760 def _get_season(self, dt): 

761 # dt = dt.date() 

762 dt = dt.replace(year=self.Y) 

763 return next(season for season, (start, end) in self.seasons if start <= dt <= end) 

764 

765 

766class TPRecDataset(KnowledgeBasedDataset): 

767 """ 

768 A dataset class for temporal recommendation tasks, inheriting from :class:`KnowledgeBasedDataset`. 

769 This class is designed only to preprocess train, valid and test sets for temporal recommendation tasks. 

770 """ 

771 

772 def __init__(self, config): 

773 super().__init__(config) 

774 

775 def build(self): 

776 datasets = super().build() 

777 # Preprocess the datasets for temporal recommendation tasks 

778 train_set = datasets[0] # train split 

779 valid_set = datasets[1] # validation split 

780 test_set = datasets[2] # test split 

781 

782 # preprocess validation and test set and link them to train so we can use it in the model as attr 

783 train_set.temporal_weights = TPRecTimestampDataset(self.config, train_set.inter_feat, "train") 

784 valid_set.temporal_weights = TPRecTimestampDataset( 

785 self.config, valid_set.inter_feat, "validation", gmm=train_set.temporal_weights.gmm_model 

786 ) 

787 test_set.temporal_weights = TPRecTimestampDataset( 

788 self.config, test_set.inter_feat, "test", gmm=train_set.temporal_weights.gmm_model 

789 ) 

790 

791 return datasets 

792 

793 

794class RPGDataset(SequentialDataset): 

795 """Dataset for :class:`~hopwise.model.sequential_recommender.rpg.RPG`. 

796 

797 Each item is mapped to a semantic ID of ``n_codebook`` digits, obtained by quantizing the item embeddings 

798 (e.g., from a sentence encoder) with OPQ. Semantic IDs are generated in :meth:`build`, after the split, 

799 so that the codebooks are trained only on the items of the training set. 

800 

801 An example of the token space when "codebook_size == 256, n_codebook == 32": 

802 0: padding 

803 1-256: digit 1 

804 257-512: digit 2 

805 ... 

806 7937-8192: digit 32 

807 8193: eos 

808 

809 Attributes: 

810 n_codebook (int): The number of digits of each semantic ID. 

811 codebook_size (int): The number of codewords of each digit. 

812 index_factory (str): The FAISS index factory string for the OPQ algorithm. 

813 item2sem_ids (dict): A dictionary mapping item ids to their semantic IDs. 

814 item2shifted_sem_id (torch.Tensor): Tensor of shape ``[item_num, n_codebook]`` mapping item ids to 

815 the tokens of their semantic IDs. 

816 eos_token (int): The end-of-sequence token. 

817 """ 

818 

819 def __init__(self, config): 

820 super().__init__(config) 

821 self.n_codebook = config["n_codebook"] 

822 self.codebook_size = config["codebook_size"] 

823 self.n_codebook_bits = self._get_codebook_bits(self.codebook_size) 

824 self.index_factory = f"OPQ{self.n_codebook},IVF1,PQ{self.n_codebook}x{self.n_codebook_bits}" 

825 self.eos_token = self.n_digit * self.codebook_size + 1 

826 

827 self.item2sem_ids = None 

828 self.item2shifted_sem_id = None 

829 

830 @property 

831 def n_digit(self): 

832 """Returns the number of digits of each semantic ID, i.e., the value of `n_codebook`.""" 

833 return self.n_codebook 

834 

835 @property 

836 def vocab_size(self) -> int: 

837 """Returns the vocabulary size, including padding and eos tokens.""" 

838 return self.eos_token + 1 

839 

840 def _get_codebook_bits(self, codebook_size): 

841 x = np.log2(codebook_size) 

842 if not x.is_integer() or x < 0: 

843 raise ValueError(f"codebook_size [{codebook_size}] should be a power of 2.") 

844 return int(x) 

845 

846 def build(self): 

847 datasets = super().build() 

848 

849 # Items seen in training, both as history and as target, are used to train the OPQ codebooks 

850 train_inter_feat = datasets[0].inter_feat 

851 training_items = torch.cat( 

852 [train_inter_feat[self.iid_field], train_inter_feat[self.item_id_list_field].flatten()] 

853 ).unique() 

854 training_items = training_items[training_items != 0].numpy() 

855 

856 item2sem_ids = self.OPQ(training_items) 

857 item2shifted_sem_id = self.shift_semantic_ids(item2sem_ids) 

858 

859 for dataset in [self, *datasets]: 

860 dataset.item2sem_ids = item2sem_ids 

861 dataset.item2shifted_sem_id = item2shifted_sem_id 

862 

863 return datasets 

864 

865 def OPQ(self, training_items): 

866 """Generates semantic IDs using the OPQ algorithm. 

867 

868 OPQ32,IVF1,PQ32x8 means: 

869 - OPQ32: Learn a rotation and split vector into 32 subspaces for quantization. 

870 - IVF1: Single inverted list (no real partitioning). 

871 - PQ32x8: Encode each of the 32 subspaces with 8 bits (256 centroids each). 

872 

873 Args: 

874 training_items (numpy.ndarray): The ids of the items used to train the index. 

875 

876 Returns: 

877 dict: A dictionary mapping item ids to their semantic IDs. 

878 """ 

879 import faiss 

880 

881 # the item embeddings are preloaded through a field that is an alias of the item id field 

882 preload_fields = [field for field in self.config["preload_weight"] or {} if field in self.alias["item_id"]] 

883 if len(preload_fields) != 1: 

884 raise ValueError( 

885 "RPG requires exactly one `preload_weight` field of item embeddings in `alias_of_item_id`, " 

886 f"found {preload_fields}." 

887 ) 

888 # rows are sorted by item id, row 0 is padding 

889 embeddings = self.get_preload_weight(preload_fields[0])[1:].astype(np.float32) 

890 

891 faiss.omp_set_num_threads(self.config["faiss_omp_num_threads"]) 

892 

893 # generate ANN from self.index_factory string using inner product to calculate distances 

894 index = faiss.index_factory(embeddings.shape[1], self.index_factory, faiss.METRIC_INNER_PRODUCT) 

895 

896 # The PQ used to learn the OPQ rotation must have the same number of bits of the index PQ 

897 opq = faiss.downcast_VectorTransform(index.chain.at(0)) 

898 custom_pq = faiss.ProductQuantizer(opq.d_out, opq.M, self.n_codebook_bits) 

899 opq.pq = custom_pq 

900 

901 self.logger.info(set_color("Training OPQ index...", "green")) 

902 index.train(embeddings[training_items - 1]) 

903 index.add(embeddings) 

904 

905 ivf_index = faiss.downcast_index(index.index) 

906 invlists = faiss.extract_index_ivf(ivf_index).invlists 

907 ls = invlists.list_size(0) 

908 # extract semantic ids, with shape |items| x |code_size| 

909 pq_codes = faiss.rev_swig_ptr(invlists.get_codes(0), ls * invlists.code_size) 

910 pq_codes = pq_codes.reshape(-1, invlists.code_size) 

911 

912 item2sem_ids = {} 

913 n_bytes = pq_codes.shape[1] 

914 for item, u8code in enumerate(pq_codes, start=1): 

915 bs = faiss.BitstringReader(faiss.swig_ptr(u8code), n_bytes) 

916 item2sem_ids[item] = tuple(bs.read(self.n_codebook_bits) for _ in range(self.n_digit)) 

917 

918 return item2sem_ids 

919 

920 def shift_semantic_ids(self, item2sem_ids): 

921 """Converts semantic IDs to tokens. 

922 

923 Each digit of a semantic ID is shifted by an offset of ``self.codebook_size * digit + 1``, 

924 such that each digit has its own token range. This is used when doing MTP (Multi Token Prediction) 

925 through multiple heads. 

926 

927 Args: 

928 item2sem_ids (dict): A dictionary mapping item ids to their semantic IDs. 

929 

930 Returns: 

931 torch.Tensor: Tensor of shape ``[item_num, n_digit]`` mapping item ids to their tokens. 

932 """ 

933 item2shifted_sem_id = torch.zeros((self.item_num, self.n_digit), dtype=torch.long) 

934 offsets = torch.arange(self.n_digit) * self.codebook_size + 1 # "+ 1" as 0 is reserved for padding 

935 for item, semantic_id_tuple in item2sem_ids.items(): 

936 item2shifted_sem_id[item] = torch.LongTensor(semantic_id_tuple) + offsets 

937 return item2shifted_sem_id 

938 

939 def __str__(self): 

940 info = [ 

941 super().__str__(), 

942 set_color("Vocabulary Size", "green") + f": {self.vocab_size}", 

943 set_color("Number of digits", "green") + f": {self.n_digit}", 

944 set_color("Codebook Size and number of PQ centroids", "green") + f": {self.codebook_size}", 

945 set_color("FAISS Configuration", "green") + f": {self.index_factory}", 

946 ] 

947 if self.item2sem_ids is not None: 

948 info.append( 

949 set_color("Percentage of unique Semantic IDs", "green") 

950 + f": {len(set(self.item2sem_ids.values())) / len(self.item2sem_ids)}" 

951 ) 

952 

953 return "\n".join(info)