Coverage for hopwise/model/sequential_recommender/s3rec.py: 92%
234 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/19 21:49
2# @Author : Hui Wang
3# @Email : hui.wang@ruc.edu.cn
5r"""S3Rec
6################################################
8Reference:
9 Kun Zhou and Hui Wang et al. "S^3-Rec: Self-Supervised Learning
10 for Sequential Recommendation with Mutual Information Maximization"
11 In CIKM 2020.
13Reference code:
14 https://github.com/RUCAIBox/CIKM2020-S3Rec
16"""
18import random
20import torch
21from torch import nn
23from hopwise.model.abstract_recommender import SequentialRecommender
24from hopwise.model.layers import TransformerEncoder
25from hopwise.model.loss import BPRLoss
28class S3Rec(SequentialRecommender):
29 r"""S3Rec is the first work to incorporate self-supervised learning in
30 sequential recommendation.
32 Note:
33 Under this framework, we need reconstruct the pretraining data,
34 which would affect the pre-training speed.
35 """
37 def __init__(self, config, dataset):
38 super().__init__(config, dataset)
40 # load parameters info
41 self.n_layers = config["n_layers"]
42 self.n_heads = config["n_heads"]
43 self.hidden_size = config["hidden_size"] # same as embedding_size
44 self.inner_size = config["inner_size"] # the dimensionality in feed-forward layer
45 self.hidden_dropout_prob = config["hidden_dropout_prob"]
46 self.attn_dropout_prob = config["attn_dropout_prob"]
47 self.hidden_act = config["hidden_act"]
48 self.layer_norm_eps = config["layer_norm_eps"]
50 self.FEATURE_FIELD = config["item_attribute"]
51 self.FEATURE_LIST = self.FEATURE_FIELD + config["LIST_SUFFIX"]
52 self.train_stage = config["train_stage"] # pretrain or finetune
53 self.pre_model_path = config["pre_model_path"] # We need this for finetune
54 self.mask_ratio = config["mask_ratio"]
55 self.aap_weight = config["aap_weight"]
56 self.mip_weight = config["mip_weight"]
57 self.map_weight = config["map_weight"]
58 self.sp_weight = config["sp_weight"]
60 self.initializer_range = config["initializer_range"]
61 self.loss_type = config["loss_type"]
63 # load dataset info
64 self.n_items = dataset.item_num + 1 # for mask token
65 self.mask_token = self.n_items - 1
66 self.n_features = dataset.num(self.FEATURE_FIELD) - 1 # we don't need padding
67 self.item_feat = dataset.get_item_feature()
69 # define layers and loss
70 # modules shared by pre-training stage and fine-tuning stage
71 self.item_embedding = nn.Embedding(self.n_items, self.hidden_size, padding_idx=0)
72 self.position_embedding = nn.Embedding(self.max_seq_length, self.hidden_size)
73 self.feature_embedding = nn.Embedding(self.n_features, self.hidden_size, padding_idx=0)
75 self.trm_encoder = TransformerEncoder(
76 n_layers=self.n_layers,
77 n_heads=self.n_heads,
78 hidden_size=self.hidden_size,
79 inner_size=self.inner_size,
80 hidden_dropout_prob=self.hidden_dropout_prob,
81 attn_dropout_prob=self.attn_dropout_prob,
82 hidden_act=self.hidden_act,
83 layer_norm_eps=self.layer_norm_eps,
84 )
86 self.LayerNorm = nn.LayerNorm(self.hidden_size, eps=self.layer_norm_eps)
87 self.dropout = nn.Dropout(self.hidden_dropout_prob)
89 # modules for pretrain
90 # add unique dense layer for 4 losses respectively
91 self.aap_norm = nn.Linear(self.hidden_size, self.hidden_size)
92 self.mip_norm = nn.Linear(self.hidden_size, self.hidden_size)
93 self.map_norm = nn.Linear(self.hidden_size, self.hidden_size)
94 self.sp_norm = nn.Linear(self.hidden_size, self.hidden_size)
95 self.loss_fct = nn.BCEWithLogitsLoss(reduction="none")
97 # modules for finetune
98 if self.loss_type == "BPR" and self.train_stage == "finetune":
99 self.loss_fct = BPRLoss()
100 elif self.loss_type == "CE" and self.train_stage == "finetune":
101 self.loss_fct = nn.CrossEntropyLoss()
102 elif self.train_stage == "finetune":
103 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!")
105 # parameters initialization
106 assert self.train_stage in ["pretrain", "finetune"]
107 if self.train_stage == "pretrain":
108 self.apply(self._init_weights)
109 else:
110 # load pretrained model for finetune
111 pretrained = torch.load(self.pre_model_path, weights_only=False)
112 self.logger.info(f"Load pretrained model from {self.pre_model_path}")
113 self.load_state_dict(pretrained["state_dict"])
115 def _init_weights(self, module):
116 """Initialize the weights"""
117 if isinstance(module, (nn.Linear, nn.Embedding)):
118 # Slightly different from the TF version which uses truncated_normal for initialization
119 # cf https://github.com/pytorch/pytorch/pull/5617
120 module.weight.data.normal_(mean=0.0, std=self.initializer_range)
121 elif isinstance(module, nn.LayerNorm):
122 module.bias.data.zero_()
123 module.weight.data.fill_(1.0)
124 if isinstance(module, nn.Linear) and module.bias is not None:
125 module.bias.data.zero_()
127 def _associated_attribute_prediction(self, sequence_output, feature_embedding):
128 sequence_output = self.aap_norm(sequence_output) # [B L H]
129 sequence_output = sequence_output.view([-1, sequence_output.size(-1), 1]) # [B*L H 1]
130 # [feature_num H] [B*L H 1] -> [B*L feature_num 1]
131 score = torch.matmul(feature_embedding, sequence_output)
132 return score.squeeze(-1) # [B*L feature_num]
134 def _masked_item_prediction(self, sequence_output, target_item_emb):
135 sequence_output = self.mip_norm(sequence_output.view([-1, sequence_output.size(-1)])) # [B*L H]
136 target_item_emb = target_item_emb.view([-1, sequence_output.size(-1)]) # [B*L H]
137 score = torch.mul(sequence_output, target_item_emb) # [B*L H]
138 return torch.sigmoid(torch.sum(score, -1)) # [B*L]
140 def _masked_attribute_prediction(self, sequence_output, feature_embedding):
141 sequence_output = self.map_norm(sequence_output) # [B L H]
142 sequence_output = sequence_output.view([-1, sequence_output.size(-1), 1]) # [B*L H 1]
143 # [feature_num H] [B*L H 1] -> [B*L feature_num 1]
144 score = torch.matmul(feature_embedding, sequence_output)
145 return score.squeeze(-1) # [B*L feature_num]
147 def _segment_prediction(self, context, segment_emb):
148 context = self.sp_norm(context)
149 score = torch.mul(context, segment_emb) # [B H]
150 return torch.sigmoid(torch.sum(score, dim=-1)) # [B]
152 def forward(self, item_seq, bidirectional=True):
153 position_ids = torch.arange(item_seq.size(1), dtype=torch.long, device=item_seq.device)
154 position_ids = position_ids.unsqueeze(0).expand_as(item_seq)
155 position_embedding = self.position_embedding(position_ids)
157 item_emb = self.item_embedding(item_seq)
158 input_emb = item_emb + position_embedding
159 input_emb = self.LayerNorm(input_emb)
160 input_emb = self.dropout(input_emb)
161 attention_mask = self.get_attention_mask(item_seq, bidirectional=bidirectional)
162 trm_output = self.trm_encoder(input_emb, attention_mask, output_all_encoded_layers=True)
163 seq_output = trm_output[-1] # [B L H]
164 return seq_output
166 def pretrain(
167 self,
168 features,
169 masked_item_sequence,
170 pos_items,
171 neg_items,
172 masked_segment_sequence,
173 pos_segment,
174 neg_segment,
175 ):
176 """Pretrain out model using four pre-training tasks:
178 1. Associated Attribute Prediction
180 2. Masked Item Prediction
182 3. Masked Attribute Prediction
184 4. Segment Prediction
185 """
186 # Encode masked sequence
187 sequence_output = self.forward(masked_item_sequence)
189 feature_embedding = self.feature_embedding.weight
190 # AAP
191 aap_score = self._associated_attribute_prediction(sequence_output, feature_embedding)
192 aap_loss = self.loss_fct(aap_score, features.view(-1, self.n_features).float())
193 # only compute loss at non-masked position
194 aap_mask = (masked_item_sequence != self.mask_token).float() * (masked_item_sequence != 0).float()
195 aap_loss = torch.sum(aap_loss * aap_mask.flatten().unsqueeze(-1))
197 # MIP
198 pos_item_embs = self.item_embedding(pos_items)
199 neg_item_embs = self.item_embedding(neg_items)
200 pos_score = self._masked_item_prediction(sequence_output, pos_item_embs)
201 neg_score = self._masked_item_prediction(sequence_output, neg_item_embs)
202 mip_distance = pos_score - neg_score
203 mip_loss = self.loss_fct(mip_distance, torch.ones_like(mip_distance, dtype=torch.float32))
204 mip_mask = (masked_item_sequence == self.mask_token).float()
205 mip_loss = torch.sum(mip_loss * mip_mask.flatten())
207 # MAP
208 map_score = self._masked_attribute_prediction(sequence_output, feature_embedding)
209 map_loss = self.loss_fct(map_score, features.view(-1, self.n_features).float())
210 map_mask = (masked_item_sequence == self.mask_token).float()
211 map_loss = torch.sum(map_loss * map_mask.flatten().unsqueeze(-1))
213 # SP
214 # segment context
215 # take the last position hidden as the context
216 segment_context = self.forward(masked_segment_sequence)[:, -1, :] # [B H]
217 pos_segment_emb = self.forward(pos_segment)[:, -1, :]
218 neg_segment_emb = self.forward(neg_segment)[:, -1, :] # [B H]
219 pos_segment_score = self._segment_prediction(segment_context, pos_segment_emb)
220 neg_segment_score = self._segment_prediction(segment_context, neg_segment_emb)
221 sp_distance = pos_segment_score - neg_segment_score
222 sp_loss = torch.sum(self.loss_fct(sp_distance, torch.ones_like(sp_distance, dtype=torch.float32)))
224 pretrain_loss = (
225 self.aap_weight * aap_loss
226 + self.mip_weight * mip_loss
227 + self.map_weight * map_loss
228 + self.sp_weight * sp_loss
229 )
231 return pretrain_loss
233 def _neg_sample(self, item_set): # [ , ]
234 item = random.randint(1, self.n_items - 1)
235 while item in item_set:
236 item = random.randint(1, self.n_items - 1)
237 return item
239 def _padding_zero_at_left(self, sequence):
240 # had truncated according to the max_length
241 pad_len = self.max_seq_length - len(sequence)
242 sequence = [0] * pad_len + sequence
243 return sequence
245 def reconstruct_pretrain_data(self, item_seq, item_seq_len):
246 """Generate pre-training data for the pre-training stage."""
247 device = item_seq.device
248 batch_size = item_seq.size(0)
250 # We don't need padding for features
251 item_feature_seq = self.item_feat[self.FEATURE_FIELD][item_seq.cpu()] - 1
253 end_index = item_seq_len.cpu().numpy().tolist()
254 item_seq = item_seq.cpu().numpy().tolist()
255 item_feature_seq = item_feature_seq.cpu().numpy().tolist()
257 # we will padding zeros at the left side
258 # these will be train_instances, after will be reshaped to batch
259 sequence_instances = []
260 associated_features = [] # For Associated Attribute Prediction and Masked Attribute Prediction
261 long_sequence = []
262 for i, end_i in enumerate(end_index):
263 sequence_instances.append(item_seq[i][:end_i])
264 long_sequence.extend(item_seq[i][:end_i])
265 # padding feature at the left side
266 associated_features.extend([[0] * self.n_features] * (self.max_seq_length - end_i))
267 for indexes in item_feature_seq[i][:end_i]:
268 features = [0] * self.n_features
269 try:
270 # multi class
271 for index in indexes:
272 if index >= 0:
273 features[index] = 1
274 except Exception:
275 # single class
276 features[indexes] = 1
277 associated_features.append(features)
279 # Masked Item Prediction and Masked Attribute Prediction
280 # [B * Len]
281 masked_item_sequence = []
282 pos_items = []
283 neg_items = []
284 for instance in sequence_instances:
285 masked_sequence = instance.copy()
286 pos_item = instance.copy()
287 neg_item = instance.copy()
288 for index_id, item in enumerate(instance):
289 prob = random.random()
290 if prob < self.mask_ratio:
291 masked_sequence[index_id] = self.mask_token
292 neg_item[index_id] = self._neg_sample(instance)
293 masked_item_sequence.append(self._padding_zero_at_left(masked_sequence))
294 pos_items.append(self._padding_zero_at_left(pos_item))
295 neg_items.append(self._padding_zero_at_left(neg_item))
297 # Segment Prediction
298 masked_segment_list = []
299 pos_segment_list = []
300 neg_segment_list = []
301 for instance in sequence_instances:
302 if len(instance) < 2: # noqa: PLR2004
303 masked_segment = instance.copy()
304 pos_segment = instance.copy()
305 neg_segment = instance.copy()
306 else:
307 sample_length = random.randint(1, len(instance) // 2)
308 start_id = random.randint(0, len(instance) - sample_length)
309 neg_start_id = random.randint(0, len(long_sequence) - sample_length)
310 pos_segment = instance[start_id : start_id + sample_length]
311 neg_segment = long_sequence[neg_start_id : neg_start_id + sample_length]
312 masked_segment = (
313 instance[:start_id] + [self.mask_token] * sample_length + instance[start_id + sample_length :]
314 )
315 pos_segment = (
316 [self.mask_token] * start_id
317 + pos_segment
318 + [self.mask_token] * (len(instance) - (start_id + sample_length))
319 )
320 neg_segment = (
321 [self.mask_token] * start_id
322 + neg_segment
323 + [self.mask_token] * (len(instance) - (start_id + sample_length))
324 )
325 masked_segment_list.append(self._padding_zero_at_left(masked_segment))
326 pos_segment_list.append(self._padding_zero_at_left(pos_segment))
327 neg_segment_list.append(self._padding_zero_at_left(neg_segment))
329 associated_features = torch.tensor(associated_features, dtype=torch.long, device=device)
330 associated_features = associated_features.view(-1, self.max_seq_length, self.n_features)
332 masked_item_sequence = torch.tensor(masked_item_sequence, dtype=torch.long, device=device).view(batch_size, -1)
333 pos_items = torch.tensor(pos_items, dtype=torch.long, device=device).view(batch_size, -1)
334 neg_items = torch.tensor(neg_items, dtype=torch.long, device=device).view(batch_size, -1)
335 masked_segment_list = torch.tensor(masked_segment_list, dtype=torch.long, device=device).view(batch_size, -1)
336 pos_segment_list = torch.tensor(pos_segment_list, dtype=torch.long, device=device).view(batch_size, -1)
337 neg_segment_list = torch.tensor(neg_segment_list, dtype=torch.long, device=device).view(batch_size, -1)
339 return (
340 associated_features,
341 masked_item_sequence,
342 pos_items,
343 neg_items,
344 masked_segment_list,
345 pos_segment_list,
346 neg_segment_list,
347 )
349 def calculate_loss(self, interaction):
350 item_seq = interaction[self.ITEM_SEQ]
351 item_seq_len = interaction[self.ITEM_SEQ_LEN]
352 # pretrain
353 if self.train_stage == "pretrain":
354 (
355 features,
356 masked_item_sequence,
357 pos_items,
358 neg_items,
359 masked_segment_sequence,
360 pos_segment,
361 neg_segment,
362 ) = self.reconstruct_pretrain_data(item_seq, item_seq_len)
364 loss = self.pretrain(
365 features,
366 masked_item_sequence,
367 pos_items,
368 neg_items,
369 masked_segment_sequence,
370 pos_segment,
371 neg_segment,
372 )
373 # finetune
374 else:
375 pos_items = interaction[self.POS_ITEM_ID]
376 # we use uni-directional attention in the fine-tuning stage
377 seq_output = self.forward(item_seq, bidirectional=False)
378 seq_output = self.gather_indexes(seq_output, item_seq_len - 1)
380 if self.loss_type == "BPR":
381 neg_items = interaction[self.NEG_ITEM_ID]
382 pos_items_emb = self.item_embedding(pos_items)
383 neg_items_emb = self.item_embedding(neg_items)
384 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) # [B]
385 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) # [B]
386 loss = self.loss_fct(pos_score, neg_score)
387 else: # self.loss_type = 'CE'
388 test_item_emb = self.item_embedding.weight
389 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1))
390 loss = self.loss_fct(logits, pos_items)
391 return loss
393 def predict(self, interaction):
394 item_seq = interaction[self.ITEM_SEQ]
395 item_seq_len = interaction[self.ITEM_SEQ_LEN]
396 test_item = interaction[self.ITEM_ID]
397 seq_output = self.forward(item_seq, bidirectional=False)
398 seq_output = self.gather_indexes(seq_output, item_seq_len - 1)
399 test_item_emb = self.item_embedding(test_item)
400 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) # [B]
401 return scores
403 def full_sort_predict(self, interaction):
404 item_seq = interaction[self.ITEM_SEQ]
405 item_seq_len = interaction[self.ITEM_SEQ_LEN]
406 seq_output = self.forward(item_seq, bidirectional=False)
407 seq_output = self.gather_indexes(seq_output, item_seq_len - 1)
408 test_items_emb = self.item_embedding.weight[: self.n_items - 1] # delete masked token
409 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) # [B, n_items]
410 return scores