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

1# @Time : 2020/9/19 21:49 

2# @Author : Hui Wang 

3# @Email : hui.wang@ruc.edu.cn 

4 

5r"""S3Rec 

6################################################ 

7 

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. 

12 

13Reference code: 

14 https://github.com/RUCAIBox/CIKM2020-S3Rec 

15 

16""" 

17 

18import random 

19 

20import torch 

21from torch import nn 

22 

23from hopwise.model.abstract_recommender import SequentialRecommender 

24from hopwise.model.layers import TransformerEncoder 

25from hopwise.model.loss import BPRLoss 

26 

27 

28class S3Rec(SequentialRecommender): 

29 r"""S3Rec is the first work to incorporate self-supervised learning in 

30 sequential recommendation. 

31 

32 Note: 

33 Under this framework, we need reconstruct the pretraining data, 

34 which would affect the pre-training speed. 

35 """ 

36 

37 def __init__(self, config, dataset): 

38 super().__init__(config, dataset) 

39 

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

49 

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

59 

60 self.initializer_range = config["initializer_range"] 

61 self.loss_type = config["loss_type"] 

62 

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

68 

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) 

74 

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 ) 

85 

86 self.LayerNorm = nn.LayerNorm(self.hidden_size, eps=self.layer_norm_eps) 

87 self.dropout = nn.Dropout(self.hidden_dropout_prob) 

88 

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

96 

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

104 

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

114 

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

126 

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] 

133 

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] 

139 

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] 

146 

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] 

151 

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) 

156 

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 

165 

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: 

177 

178 1. Associated Attribute Prediction 

179 

180 2. Masked Item Prediction 

181 

182 3. Masked Attribute Prediction 

183 

184 4. Segment Prediction 

185 """ 

186 # Encode masked sequence 

187 sequence_output = self.forward(masked_item_sequence) 

188 

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

196 

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

206 

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

212 

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

223 

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 ) 

230 

231 return pretrain_loss 

232 

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 

238 

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 

244 

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) 

249 

250 # We don't need padding for features 

251 item_feature_seq = self.item_feat[self.FEATURE_FIELD][item_seq.cpu()] - 1 

252 

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

256 

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) 

278 

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

296 

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

328 

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) 

331 

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) 

338 

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 ) 

348 

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) 

363 

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) 

379 

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 

392 

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 

402 

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