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
« 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
5# UPDATE
6# @Time : 2021/7/9
7# @Author : Yupeng Hou
8# @Email : houyupeng@ruc.edu.cn
10"""hopwise.data.customized_dataset
11##################################
13We only recommend building customized datasets by inheriting.
15Customized datasets named ``[Model Name]Dataset`` can be automatically called.
16"""
18import datetime
20import numba
21import numpy as np
22import pandas as pd
23import torch
24from sklearn.mixture import GaussianMixture as GMM
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
39class GRU4RecKGDataset(KGSeqDataset):
40 def __init__(self, config):
41 super().__init__(config)
44class KSRDataset(KGSeqDataset):
45 def __init__(self, config):
46 super().__init__(config)
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.
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.
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 """
65 def __init__(self, config):
66 super().__init__(config)
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])
74 def data_augmentation(self):
75 """Augmentation processing for sequential dataset.
77 E.g., ``u1`` has purchase sequence ``<i1, i2, i3, i4>``,
78 then after augmentation, we will generate three cases.
80 ``u1, <i1> | i2``
82 (Which means given user_id ``u1`` and item_seq ``<i1>``,
83 we need to predict the next item ``i2``.)
85 The other cases are below:
87 ``u1, <i1, i2> | i3``
89 ``u1, <i1, i2, i3> | i4``
90 """
91 self.logger.debug("data_augmentation")
93 self._aug_presets()
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)
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)
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 }
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)
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]
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]
148 new_data.update(Interaction(new_dict))
149 self.inter_feat = new_data
152class KGGLMDatasetMixin:
153 """Mixin class containing KGGLM-specific dataset logic.
155 This mixin should be used with KnowledgePathDataset or UserItemKnowledgePathDataset
156 to create the appropriate KGGLM dataset class.
157 """
159 def _get_field_from_config(self):
160 super()._get_field_from_config()
161 self.train_stage = self.config["train_stage"]
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"]
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()
175 def generate_pretrain_dataset(self):
176 """Generate pretrain dataset for KGGLM model using CSR-based parallel random walks."""
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()
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)
187 min_hop, max_hop = self.pretrain_hop_length
188 max_tries_per_entity = self.config["path_sample_args"]["MAX_RW_TRIES_PER_IID"]
190 entity_ids = self._get_entity_ids_range()
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)
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))
199 self.logger.info(
200 set_color(f"Running {len(all_start_nodes)} parallel random walks for pretraining...", "blue")
201 )
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 )
208 # Deduplicate and collect unique paths per entity
209 unique_paths = set()
210 entity_path_counts = {}
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 )
219 for idx in iter_paths:
220 entity = all_start_nodes[idx]
221 hop_length = all_hop_lengths[idx]
223 # Check if entity already has enough paths
224 if entity_path_counts.get(entity, 0) >= self.pretrain_paths:
225 continue
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])
231 # Skip if path has invalid nodes (walk got stuck)
232 if -1 in path:
233 continue
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)
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
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)
253 self._path_dataset = path_string
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.
260 Note: Set np.random.seed() before calling this function for reproducibility.
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)
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)
279 for i in numba.prange(n_walks):
280 node = start_nodes[i]
281 num_steps = hop_lengths[i]
282 paths[i, 0] = node
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
289 if n_neighbors == 0:
290 break
292 # random choice replacement for numba with p=weights
293 edge_weights = weights[start_idx:end_idx]
294 total_weight = edge_weights.sum()
296 if total_weight == 0:
297 break
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
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]
313 return paths, path_relations
316class KGGLMDataset(KGGLMDatasetMixin, KnowledgePathDataset):
317 """KGGLM dataset inheriting from KnowledgePathDataset."""
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)
328class UserItemKGGLMDataset(KGGLMDatasetMixin, UserItemKnowledgePathDataset):
329 """KGGLM dataset inheriting from UserItemKnowledgePathDataset.
331 Used when both user and item knowledge graph links are available.
332 """
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
351class TPRecTimestampDataset:
352 """
353 A class to create a clustered dataset based on interaction timestamp specifically for the TPRec model
354 """
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'
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
374 self.user_item_timestamp = np.array(self.data.timestamps)
375 time2num = self._timeanalysis()
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 )
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"])
397 ucp_hash = self._generate_clus_dict(labels)
398 uc_weight = self._generate_user_agent_num(ucp_hash)
400 self.uc_weight = uc_weight
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 )
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)
478 return test_cluster_label
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 )
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])
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)
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)
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()
564 return fileTime
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])
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)
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)
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()
602 return fileTime
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 )
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 )
656 return fileTime
658 def _timeanalysis(self):
659 dac_time = self.data.timestamps.value_counts()
660 dac_time_date = pd.to_datetime(dac_time.index)
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]
680 for gap in [90, 30, 7, 1]:
681 self._structuralWithGap(gap, time2num, time2relative, mapIndex, serial2PrefixSum)
682 return time2num
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))
689 models = GMM(self.config["cluster_num"], covariance_type="full", random_state=self.config["seed"]).fit(x)
690 labels = models.predict(x)
692 for user, item, label in zip(self.data.users, self.data.item, labels):
693 ui2label[(user, item)] = label
695 return models, self.config["cluster_num"], labels, ui2label
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
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)
713 return ucp_hash
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
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
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
757 time2Num[time2relative[idx]].append(gap_left)
758 time2Num[time2relative[idx]].append(gap_left_2)
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)
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 """
772 def __init__(self, config):
773 super().__init__(config)
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
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 )
791 return datasets
794class RPGDataset(SequentialDataset):
795 """Dataset for :class:`~hopwise.model.sequential_recommender.rpg.RPG`.
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.
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
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 """
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
827 self.item2sem_ids = None
828 self.item2shifted_sem_id = None
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
835 @property
836 def vocab_size(self) -> int:
837 """Returns the vocabulary size, including padding and eos tokens."""
838 return self.eos_token + 1
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)
846 def build(self):
847 datasets = super().build()
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()
856 item2sem_ids = self.OPQ(training_items)
857 item2shifted_sem_id = self.shift_semantic_ids(item2sem_ids)
859 for dataset in [self, *datasets]:
860 dataset.item2sem_ids = item2sem_ids
861 dataset.item2shifted_sem_id = item2shifted_sem_id
863 return datasets
865 def OPQ(self, training_items):
866 """Generates semantic IDs using the OPQ algorithm.
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).
873 Args:
874 training_items (numpy.ndarray): The ids of the items used to train the index.
876 Returns:
877 dict: A dictionary mapping item ids to their semantic IDs.
878 """
879 import faiss
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)
891 faiss.omp_set_num_threads(self.config["faiss_omp_num_threads"])
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)
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
901 self.logger.info(set_color("Training OPQ index...", "green"))
902 index.train(embeddings[training_items - 1])
903 index.add(embeddings)
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)
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))
918 return item2sem_ids
920 def shift_semantic_ids(self, item2sem_ids):
921 """Converts semantic IDs to tokens.
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.
927 Args:
928 item2sem_ids (dict): A dictionary mapping item ids to their semantic IDs.
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
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 )
953 return "\n".join(info)