Coverage for hopwise/data/transform.py: 98%

180 statements  

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

1# @Time : 2022/7/19 

2# @Author : Gaowei Zhang 

3# @Email : zgw15630559577@163.com 

4 

5import math 

6import random 

7from copy import deepcopy 

8 

9import numpy as np 

10import torch 

11 

12from hopwise.data.interaction import Interaction 

13 

14 

15def construct_transform(config): 

16 """Transformation for batch data.""" 

17 if config["transform"] is None: 

18 return Equal(config) 

19 else: 

20 str2transform = { 

21 "mask_itemseq": MaskItemSequence, 

22 "inverse_itemseq": InverseItemSequence, 

23 "crop_itemseq": CropItemSequence, 

24 "reorder_itemseq": ReorderItemSequence, 

25 "user_defined": UserDefinedTransform, 

26 } 

27 if config["transform"] not in str2transform: 

28 raise NotImplementedError(f"There is no transform named '{config['transform']}'") 

29 

30 return str2transform[config["transform"]](config) 

31 

32 

33class Equal: 

34 def __init__(self, config): 

35 pass 

36 

37 def __call__(self, dataset, interaction): 

38 return interaction 

39 

40 

41class MaskItemSequence: 

42 """Mask item sequence for training.""" 

43 

44 def __init__(self, config): 

45 self.ITEM_SEQ = config["ITEM_ID_FIELD"] + config["LIST_SUFFIX"] 

46 self.ITEM_ID = config["ITEM_ID_FIELD"] 

47 self.MASK_ITEM_SEQ = "Mask_" + self.ITEM_SEQ 

48 self.POS_ITEMS = "Pos_" + config["ITEM_ID_FIELD"] 

49 self.NEG_ITEMS = "Neg_" + config["ITEM_ID_FIELD"] 

50 self.max_seq_length = config["MAX_ITEM_LIST_LENGTH"] 

51 self.mask_ratio = config["mask_ratio"] 

52 self.ft_ratio = 0 if not hasattr(config, "ft_ratio") else config["ft_ratio"] 

53 self.mask_item_length = int(self.mask_ratio * self.max_seq_length) 

54 self.MASK_INDEX = "MASK_INDEX" 

55 config["MASK_INDEX"] = "MASK_INDEX" 

56 config["MASK_ITEM_SEQ"] = self.MASK_ITEM_SEQ 

57 config["POS_ITEMS"] = self.POS_ITEMS 

58 config["NEG_ITEMS"] = self.NEG_ITEMS 

59 self.ITEM_SEQ_LEN = config["ITEM_LIST_LENGTH_FIELD"] 

60 self.config = config 

61 

62 def _neg_sample(self, item_set, n_items): 

63 item = random.randint(1, n_items - 1) 

64 while item in item_set: 

65 item = random.randint(1, n_items - 1) 

66 return item 

67 

68 def _padding_sequence(self, sequence, max_length): 

69 pad_len = max_length - len(sequence) 

70 sequence = [0] * pad_len + sequence 

71 sequence = sequence[-max_length:] # truncate according to the max_length 

72 return sequence 

73 

74 def _append_mask_last(self, interaction, n_items, device): 

75 batch_size = interaction[self.ITEM_SEQ].size(0) 

76 pos_items, neg_items, masked_index, masked_item_sequence = [], [], [], [] 

77 seq_instance = interaction[self.ITEM_SEQ].cpu().numpy().tolist() 

78 item_seq_len = interaction[self.ITEM_SEQ_LEN].cpu().numpy().tolist() 

79 for instance, lens in zip(seq_instance, item_seq_len): 

80 mask_seq = instance.copy() 

81 ext = instance[lens - 1] 

82 mask_seq[lens - 1] = n_items 

83 masked_item_sequence.append(mask_seq) 

84 pos_items.append(self._padding_sequence([ext], self.mask_item_length)) 

85 neg_items.append(self._padding_sequence([self._neg_sample(instance, n_items)], self.mask_item_length)) 

86 masked_index.append(self._padding_sequence([lens - 1], self.mask_item_length)) 

87 # [B Len] 

88 masked_item_sequence = torch.tensor(masked_item_sequence, dtype=torch.long, device=device).view(batch_size, -1) 

89 # [B mask_len] 

90 pos_items = torch.tensor(pos_items, dtype=torch.long, device=device).view(batch_size, -1) 

91 # [B mask_len] 

92 neg_items = torch.tensor(neg_items, dtype=torch.long, device=device).view(batch_size, -1) 

93 # [B mask_len] 

94 masked_index = torch.tensor(masked_index, dtype=torch.long, device=device).view(batch_size, -1) 

95 new_dict = { 

96 self.MASK_ITEM_SEQ: masked_item_sequence, 

97 self.POS_ITEMS: pos_items, 

98 self.NEG_ITEMS: neg_items, 

99 self.MASK_INDEX: masked_index, 

100 } 

101 ft_interaction = deepcopy(interaction) 

102 ft_interaction.update(Interaction(new_dict)) 

103 return ft_interaction 

104 

105 def __call__(self, dataset, interaction): 

106 item_seq = interaction[self.ITEM_SEQ] 

107 device = item_seq.device 

108 batch_size = item_seq.size(0) 

109 n_items = dataset.num(self.ITEM_ID) 

110 sequence_instances = item_seq.cpu().numpy().tolist() 

111 

112 # Masked Item Prediction 

113 # [B * Len] 

114 masked_item_sequence = [] 

115 pos_items = [] 

116 neg_items = [] 

117 masked_index = [] 

118 

119 if random.random() < self.ft_ratio: 

120 interaction = self._append_mask_last(interaction, n_items, device) 

121 else: 

122 for instance in sequence_instances: 

123 # WE MUST USE 'copy()' HERE! 

124 masked_sequence = instance.copy() 

125 pos_item = [] 

126 neg_item = [] 

127 index_ids = [] 

128 for index_id, item in enumerate(instance): 

129 # padding is 0, the sequence is end 

130 if item == 0: 

131 break 

132 prob = random.random() 

133 if prob < self.mask_ratio: 

134 pos_item.append(item) 

135 neg_item.append(self._neg_sample(instance, n_items)) 

136 masked_sequence[index_id] = n_items 

137 index_ids.append(index_id) 

138 

139 masked_item_sequence.append(masked_sequence) 

140 pos_items.append(self._padding_sequence(pos_item, self.mask_item_length)) 

141 neg_items.append(self._padding_sequence(neg_item, self.mask_item_length)) 

142 masked_index.append(self._padding_sequence(index_ids, self.mask_item_length)) 

143 

144 # [B Len] 

145 masked_item_sequence = torch.tensor(masked_item_sequence, dtype=torch.long, device=device).view( 

146 batch_size, -1 

147 ) 

148 # [B mask_len] 

149 pos_items = torch.tensor(pos_items, dtype=torch.long, device=device).view(batch_size, -1) 

150 # [B mask_len] 

151 neg_items = torch.tensor(neg_items, dtype=torch.long, device=device).view(batch_size, -1) 

152 # [B mask_len] 

153 masked_index = torch.tensor(masked_index, dtype=torch.long, device=device).view(batch_size, -1) 

154 new_dict = { 

155 self.MASK_ITEM_SEQ: masked_item_sequence, 

156 self.POS_ITEMS: pos_items, 

157 self.NEG_ITEMS: neg_items, 

158 self.MASK_INDEX: masked_index, 

159 } 

160 interaction.update(Interaction(new_dict)) 

161 return interaction 

162 

163 

164class InverseItemSequence: 

165 """inverse the seq_item, like this 

166 [1,2,3,0,0,0,0] -- after inverse -->> [0,0,0,0,1,2,3] 

167 """ 

168 

169 def __init__(self, config): 

170 self.ITEM_SEQ = config["ITEM_ID_FIELD"] + config["LIST_SUFFIX"] 

171 self.ITEM_SEQ_LEN = config["ITEM_LIST_LENGTH_FIELD"] 

172 self.INVERSE_ITEM_SEQ = "Inverse_" + self.ITEM_SEQ 

173 config["INVERSE_ITEM_SEQ"] = self.INVERSE_ITEM_SEQ 

174 

175 def __call__(self, dataset, interaction): 

176 item_seq = interaction[self.ITEM_SEQ] 

177 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

178 device = item_seq.device 

179 item_seq = item_seq.cpu().numpy() 

180 item_seq_len = item_seq_len.cpu().numpy() 

181 new_item_seq = [] 

182 for items, length in zip(item_seq, item_seq_len): 

183 item = list(items[:length]) 

184 zeros = list(items[length:]) 

185 seqs = zeros + item 

186 new_item_seq.append(seqs) 

187 inverse_item_seq = torch.tensor(new_item_seq, dtype=torch.long, device=device) 

188 new_dict = {self.INVERSE_ITEM_SEQ: inverse_item_seq} 

189 interaction.update(Interaction(new_dict)) 

190 return interaction 

191 

192 

193class CropItemSequence: 

194 """Random crop for item sequence.""" 

195 

196 def __init__(self, config): 

197 self.ITEM_SEQ = config["ITEM_ID_FIELD"] + config["LIST_SUFFIX"] 

198 self.CROP_ITEM_SEQ = "Crop_" + self.ITEM_SEQ 

199 self.ITEM_SEQ_LEN = config["ITEM_LIST_LENGTH_FIELD"] 

200 self.CROP_ITEM_SEQ_LEN = self.CROP_ITEM_SEQ + self.ITEM_SEQ_LEN 

201 self.crop_eta = config["eta"] 

202 config["CROP_ITEM_SEQ"] = self.CROP_ITEM_SEQ 

203 config["CROP_ITEM_SEQ_LEN"] = self.CROP_ITEM_SEQ_LEN 

204 

205 def __call__(self, dataset, interaction): 

206 item_seq = interaction[self.ITEM_SEQ] 

207 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

208 device = item_seq.device 

209 crop_item_seq_list, crop_item_seqlen_list = [], [] 

210 

211 for seq, length in zip(item_seq, item_seq_len): 

212 crop_len = math.floor(length * self.crop_eta) 

213 crop_begin = random.randint(0, length - crop_len) 

214 crop_item_seq = np.zeros(seq.shape[0]) 

215 if crop_begin + crop_len < seq.shape[0]: 

216 crop_item_seq[:crop_len] = seq[crop_begin : crop_begin + crop_len] 

217 else: 

218 crop_item_seq[:crop_len] = seq[crop_begin:] 

219 crop_item_seq_list.append(torch.tensor(crop_item_seq, dtype=torch.long, device=device)) 

220 crop_item_seqlen_list.append(torch.tensor(crop_len, dtype=torch.long, device=device)) 

221 new_dict = { 

222 self.CROP_ITEM_SEQ: torch.stack(crop_item_seq_list), 

223 self.CROP_ITEM_SEQ_LEN: torch.stack(crop_item_seqlen_list), 

224 } 

225 interaction.update(Interaction(new_dict)) 

226 return interaction 

227 

228 

229class ReorderItemSequence: 

230 """Reorder operation for item sequence.""" 

231 

232 def __init__(self, config): 

233 self.ITEM_SEQ = config["ITEM_ID_FIELD"] + config["LIST_SUFFIX"] 

234 self.REORDER_ITEM_SEQ = "Reorder_" + self.ITEM_SEQ 

235 self.ITEM_SEQ_LEN = config["ITEM_LIST_LENGTH_FIELD"] 

236 self.reorder_beta = config["beta"] 

237 config["REORDER_ITEM_SEQ"] = self.REORDER_ITEM_SEQ 

238 

239 def __call__(self, dataset, interaction): 

240 item_seq = interaction[self.ITEM_SEQ] 

241 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

242 device = item_seq.device 

243 reorder_seq_list = [] 

244 

245 for seq, length in zip(item_seq, item_seq_len): 

246 reorder_len = math.floor(length * self.reorder_beta) 

247 reorder_begin = random.randint(0, length - reorder_len) 

248 reorder_item_seq = seq.cpu().detach().numpy().copy() 

249 

250 shuffle_index = list(range(reorder_begin, reorder_begin + reorder_len)) 

251 random.shuffle(shuffle_index) 

252 reorder_item_seq[reorder_begin : reorder_begin + reorder_len] = reorder_item_seq[shuffle_index] 

253 

254 reorder_seq_list.append(torch.tensor(reorder_item_seq, dtype=torch.long, device=device)) 

255 new_dict = {self.REORDER_ITEM_SEQ: torch.stack(reorder_seq_list)} 

256 interaction.update(Interaction(new_dict)) 

257 return interaction 

258 

259 

260class UserDefinedTransform: 

261 def __init__(self, config): 

262 pass 

263 

264 def __call__(self, dataset, interaction): 

265 pass