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
« 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
5import math
6import random
7from copy import deepcopy
9import numpy as np
10import torch
12from hopwise.data.interaction import Interaction
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']}'")
30 return str2transform[config["transform"]](config)
33class Equal:
34 def __init__(self, config):
35 pass
37 def __call__(self, dataset, interaction):
38 return interaction
41class MaskItemSequence:
42 """Mask item sequence for training."""
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
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
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
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
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()
112 # Masked Item Prediction
113 # [B * Len]
114 masked_item_sequence = []
115 pos_items = []
116 neg_items = []
117 masked_index = []
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)
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))
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
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 """
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
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
193class CropItemSequence:
194 """Random crop for item sequence."""
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
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 = [], []
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
229class ReorderItemSequence:
230 """Reorder operation for item sequence."""
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
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 = []
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()
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]
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
260class UserDefinedTransform:
261 def __init__(self, config):
262 pass
264 def __call__(self, dataset, interaction):
265 pass