Coverage for hopwise/data/dataset/sequential_dataset.py: 96%
101 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/9/16
2# @Author : Yushuo Chen
3# @Email : chenyushuo@ruc.edu.cn
5# UPDATE:
6# @Time : 2022/7/8, 2020/9/16, 2021/7/1, 2021/7/11
7# @Author : Zhen Tian, Yushuo Chen, Xingyu Pan, Yupeng Hou
8# @Email : chenyuwuxinn@gmail.com, chenyushuo@ruc.edu.cn, xy_pan@foxmail.com, houyupeng@ruc.edu.cn
10"""hopwise.data.sequential_dataset
11###############################
12"""
14import numpy as np
15import torch
17from hopwise.data.dataset import Dataset
18from hopwise.data.interaction import Interaction
19from hopwise.utils import FeatureSource, FeatureType
22class SequentialDataset(Dataset):
23 """:class:`SequentialDataset` is based on :class:`~hopwise.data.dataset.dataset.Dataset`,
24 and provides augmentation interface to adapt to Sequential Recommendation,
25 which can accelerate the data loader.
27 Attributes:
28 max_item_list_len (int): Max length of historical item list.
29 item_list_length_field (str): Field name for item lists' length.
30 """
32 def __init__(self, config):
33 self.max_item_list_len = config["MAX_ITEM_LIST_LENGTH"]
34 self.item_list_length_field = config["ITEM_LIST_LENGTH_FIELD"]
35 super().__init__(config)
36 if config["benchmark_filename"] is not None:
37 self._benchmark_presets()
39 def _change_feat_format(self):
40 """Change feat format from :class:`pandas.DataFrame` to :class:`Interaction`,
41 then perform data augmentation.
42 """
43 super()._change_feat_format()
45 if self.config["benchmark_filename"] is not None:
46 return
47 self.logger.debug("Augmentation for sequential recommendation.")
48 self.data_augmentation()
50 def _aug_presets(self):
51 list_suffix = self.config["LIST_SUFFIX"]
52 for field in self.inter_feat:
53 if field != self.uid_field:
54 list_field = field + list_suffix
55 setattr(self, f"{field}_list_field", list_field)
56 ftype = self.field2type[field]
58 if ftype in [FeatureType.TOKEN, FeatureType.TOKEN_SEQ]:
59 list_ftype = FeatureType.TOKEN_SEQ
60 else:
61 list_ftype = FeatureType.FLOAT_SEQ
63 if ftype in [FeatureType.TOKEN_SEQ, FeatureType.FLOAT_SEQ]:
64 list_len = (self.max_item_list_len, self.field2seqlen[field])
65 else:
66 list_len = self.max_item_list_len
68 self.set_field_property(list_field, list_ftype, FeatureSource.INTERACTION, list_len)
70 self.set_field_property(self.item_list_length_field, FeatureType.TOKEN, FeatureSource.INTERACTION, 1)
72 def data_augmentation(self):
73 """Augmentation processing for sequential dataset.
75 E.g., ``u1`` has purchase sequence ``<i1, i2, i3, i4>``,
76 then after augmentation, we will generate three cases.
78 ``u1, <i1> | i2``
80 (Which means given user_id ``u1`` and item_seq ``<i1>``,
81 we need to predict the next item ``i2``.)
83 The other cases are below:
85 ``u1, <i1, i2> | i3``
87 ``u1, <i1, i2, i3> | i4``
88 """
89 self.logger.debug("data_augmentation")
91 self._aug_presets()
93 self._check_field("uid_field", "time_field")
94 max_item_list_len = self.config["MAX_ITEM_LIST_LENGTH"]
95 self.sort(by=[self.uid_field, self.time_field], ascending=True)
96 last_uid = None
97 uid_list, item_list_index, target_index, item_list_length = [], [], [], []
98 seq_start = 0
99 for i, uid in enumerate(self.inter_feat[self.uid_field].numpy()):
100 if last_uid != uid:
101 last_uid = uid
102 seq_start = i
103 else:
104 if i - seq_start > max_item_list_len:
105 seq_start += 1
106 uid_list.append(uid)
107 item_list_index.append(slice(seq_start, i))
108 target_index.append(i)
109 item_list_length.append(i - seq_start)
111 uid_list = np.array(uid_list)
112 item_list_index = np.array(item_list_index)
113 target_index = np.array(target_index)
114 item_list_length = np.array(item_list_length, dtype=np.int64)
116 new_length = len(item_list_index)
117 new_data = self.inter_feat[target_index]
118 new_dict = {
119 self.item_list_length_field: torch.tensor(item_list_length),
120 }
122 for field in self.inter_feat:
123 if field != self.uid_field:
124 list_field = getattr(self, f"{field}_list_field")
125 list_len = self.field2seqlen[list_field]
126 shape = (new_length, list_len) if isinstance(list_len, int) else (new_length,) + list_len
127 if (
128 self.field2type[field] in [FeatureType.FLOAT, FeatureType.FLOAT_SEQ]
129 and field in self.config["numerical_features"]
130 ):
131 shape += (2,)
132 new_dict[list_field] = torch.zeros(shape, dtype=self.inter_feat[field].dtype)
134 value = self.inter_feat[field]
135 for i, (index, length) in enumerate(zip(item_list_index, item_list_length)):
136 new_dict[list_field][i][:length] = value[index]
138 new_data.update(Interaction(new_dict))
139 self.inter_feat = new_data
141 def _benchmark_presets(self):
142 list_suffix = self.config["LIST_SUFFIX"]
143 for field in self.inter_feat:
144 if field + list_suffix in self.inter_feat:
145 list_field = field + list_suffix
146 setattr(self, f"{field}_list_field", list_field)
147 self.set_field_property(self.item_list_length_field, FeatureType.TOKEN, FeatureSource.INTERACTION, 1)
148 self.inter_feat[self.item_list_length_field] = self.inter_feat[self.item_id_list_field].transform(len)
150 def inter_matrix(self, form="coo", value_field=None):
151 """Get sparse matrix that describe interactions between user_id and item_id.
152 Sparse matrix has shape (user_num, item_num).
153 For a row of <src, tgt>, ``matrix[src, tgt] = 1`` if ``value_field`` is ``None``,
154 else ``matrix[src, tgt] = self.inter_feat[src, tgt]``.
156 Args:
157 form (str, optional): Sparse matrix format. Defaults to ``coo``.
158 value_field (str, optional): Data of sparse matrix, which should exist in ``df_feat``.
159 Defaults to ``None``.
161 Returns:
162 scipy.sparse: Sparse matrix in form ``coo`` or ``csr``.
163 """
164 if not self.uid_field or not self.iid_field:
165 raise ValueError("dataset does not exist uid/iid, thus can not converted to sparse matrix.")
167 l1_idx = self.inter_feat[self.item_list_length_field] == 1
168 l1_inter_dict = self.inter_feat[l1_idx].interaction
169 new_dict = {}
170 list_suffix = self.config["LIST_SUFFIX"]
171 candidate_field_set = set()
172 for field in l1_inter_dict:
173 if field != self.uid_field and field + list_suffix in l1_inter_dict:
174 candidate_field_set.add(field)
175 new_dict[field] = torch.cat([self.inter_feat[field], l1_inter_dict[field + list_suffix][:, 0]])
176 elif (not field.endswith(list_suffix)) and (field != self.item_list_length_field):
177 new_dict[field] = torch.cat([self.inter_feat[field], l1_inter_dict[field]])
178 local_inter_feat = Interaction(new_dict)
179 return self._create_sparse_matrix(local_inter_feat, self.uid_field, self.iid_field, form, value_field)
181 def build(self):
182 """Processing dataset according to evaluation setting, including Group, Order and Split.
183 See :class:`~hopwise.config.eval_setting.EvalSetting` for details.
185 Args:
186 eval_setting (:class:`~hopwise.config.eval_setting.EvalSetting`):
187 Object contains evaluation settings, which guide the data processing procedure.
189 Returns:
190 list: List of built :class:`Dataset`.
191 """
192 ordering_args = self.config["eval_args"]["order"]
193 if ordering_args != "TO":
194 raise ValueError("The ordering args for sequential recommendation has to be 'TO'")
196 return super().build()