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

1# @Time : 2020/9/16 

2# @Author : Yushuo Chen 

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

4 

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 

9 

10"""hopwise.data.sequential_dataset 

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

12""" 

13 

14import numpy as np 

15import torch 

16 

17from hopwise.data.dataset import Dataset 

18from hopwise.data.interaction import Interaction 

19from hopwise.utils import FeatureSource, FeatureType 

20 

21 

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. 

26 

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

31 

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

38 

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

44 

45 if self.config["benchmark_filename"] is not None: 

46 return 

47 self.logger.debug("Augmentation for sequential recommendation.") 

48 self.data_augmentation() 

49 

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] 

57 

58 if ftype in [FeatureType.TOKEN, FeatureType.TOKEN_SEQ]: 

59 list_ftype = FeatureType.TOKEN_SEQ 

60 else: 

61 list_ftype = FeatureType.FLOAT_SEQ 

62 

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 

67 

68 self.set_field_property(list_field, list_ftype, FeatureSource.INTERACTION, list_len) 

69 

70 self.set_field_property(self.item_list_length_field, FeatureType.TOKEN, FeatureSource.INTERACTION, 1) 

71 

72 def data_augmentation(self): 

73 """Augmentation processing for sequential dataset. 

74 

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

76 then after augmentation, we will generate three cases. 

77 

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

79 

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

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

82 

83 The other cases are below: 

84 

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

86 

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

88 """ 

89 self.logger.debug("data_augmentation") 

90 

91 self._aug_presets() 

92 

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) 

110 

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) 

115 

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 } 

121 

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) 

133 

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] 

137 

138 new_data.update(Interaction(new_dict)) 

139 self.inter_feat = new_data 

140 

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) 

149 

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

155 

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

160 

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

166 

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) 

180 

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. 

184 

185 Args: 

186 eval_setting (:class:`~hopwise.config.eval_setting.EvalSetting`): 

187 Object contains evaluation settings, which guide the data processing procedure. 

188 

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

195 

196 return super().build()