Coverage for hopwise/model/sequential_recommender/sasreccpr.py: 83%

212 statements  

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

1# @Time : 2020/9/18 11:33 

2# @Author : Hui Wang 

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

4 

5# UPDATE: 

6# @Time : 2023/11/24 

7# @Author : Haw-Shiuan Chang 

8# @Email : ken77921@gmail.com 

9 

10"""SASRec + Softmax-CPR 

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

12 

13Reference: 

14 Wang-Cheng Kang et al. "Self-Attentive Sequential Recommendation." in ICDM 2018. 

15 Haw-Shiuan Chang, Nikhil Agarwal, and Andrew McCallum "To Copy, or not to Copy; That is a Critical Issue of the Output Softmax Layer in Neural Sequential Recommenders" in WSDM 2024 

16 

17Reference: 

18 https://github.com/kang205/SASRec 

19 https://arxiv.org/pdf/2310.14079.pdf 

20 

21""" # noqa: E501 

22 

23# from hopwise.model.loss import BPRLoss 

24import math 

25import sys 

26 

27import torch 

28import torch.nn.functional as F 

29from torch import nn 

30 

31from hopwise.model.abstract_recommender import SequentialRecommender 

32from hopwise.model.layers import TransformerEncoder 

33 

34 

35def gelu(x): 

36 return 0.5 * x * (1 + torch.tanh(math.sqrt(2 / math.pi) * (x + 0.044715 * torch.pow(x, 3)))) 

37 

38 

39class SASRecCPR(SequentialRecommender): 

40 r"""SASRec is the first sequential recommender based on self-attentive mechanism. 

41 

42 Note: 

43 In the author's implementation, the Point-Wise Feed-Forward Network (PFFN) is implemented 

44 by CNN with 1x1 kernel. In this implementation, we follows the original BERT implementation 

45 using Fully Connected Layer to implement the PFFN. 

46 """ 

47 

48 def __init__(self, config, dataset): 

49 super().__init__(config, dataset) 

50 

51 # load parameters info 

52 self.n_layers = config["n_layers"] 

53 self.n_heads = config["n_heads"] 

54 self.hidden_size = config["hidden_size"] # same as embedding_size 

55 self.inner_size = config["inner_size"] # the dimensionality in feed-forward layer 

56 self.hidden_dropout_prob = config["hidden_dropout_prob"] 

57 self.attn_dropout_prob = config["attn_dropout_prob"] 

58 self.hidden_act = config["hidden_act"] 

59 self.layer_norm_eps = config["layer_norm_eps"] 

60 self.initializer_range = config["initializer_range"] 

61 self.loss_type = config["loss_type"] 

62 self.n_facet_all = config["n_facet_all"] # added for mfs 

63 self.n_facet = config["n_facet"] # added for mfs 

64 self.n_facet_window = config["n_facet_window"] # added for mfs 

65 self.n_facet_hidden = min(config["n_facet_hidden"], config["n_layers"]) # added for mfs 

66 self.n_facet_MLP = config["n_facet_MLP"] # added for mfs 

67 self.n_facet_context = config["n_facet_context"] # added for dynamic partioning 

68 self.n_facet_reranker = config["n_facet_reranker"] # added for dynamic partioning 

69 self.n_facet_emb = config["n_facet_emb"] # added for dynamic partioning 

70 self.weight_mode = config["weight_mode"] # added for mfs 

71 self.context_norm = config["context_norm"] # added for mfs 

72 self.post_remove_context = config["post_remove_context"] # added for mfs 

73 self.partition_merging_mode = config["partition_merging_mode"] # added for mfs 

74 self.reranker_merging_mode = config["reranker_merging_mode"] # added for mfs 

75 self.reranker_CAN_NUM = [int(x) for x in str(config["reranker_CAN_NUM"]).split(",")] 

76 self.candidates_from_previous_reranker = True 

77 if self.weight_mode == "max_logits": 

78 self.n_facet_effective = 1 

79 else: 

80 self.n_facet_effective = self.n_facet 

81 

82 assert ( 

83 self.n_facet + self.n_facet_context + self.n_facet_reranker * len(self.reranker_CAN_NUM) + self.n_facet_emb 

84 == self.n_facet_all 

85 ) 

86 assert self.n_facet_emb in (0, 2) 

87 assert self.n_facet_MLP <= 0 # -1 or 0 

88 assert self.n_facet_window <= 0 

89 self.n_facet_window = -self.n_facet_window 

90 self.n_facet_MLP = -self.n_facet_MLP 

91 self.softmax_nonlinear = "None" # added for mfs 

92 self.use_proj_bias = config["use_proj_bias"] # added for mfs 

93 hidden_state_input_ratio = 1 + self.n_facet_MLP # 1 + 1 

94 self.MLP_linear = nn.Linear( 

95 self.hidden_size * (self.n_facet_hidden * (self.n_facet_window + 1)), 

96 self.hidden_size * self.n_facet_MLP, 

97 ) # (hid_dim*2) -> (hid_dim) 

98 total_lin_dim = self.hidden_size * hidden_state_input_ratio 

99 self.project_arr = nn.ModuleList( 

100 [nn.Linear(total_lin_dim, self.hidden_size, bias=self.use_proj_bias) for i in range(self.n_facet_all)] 

101 ) 

102 

103 self.project_emb = nn.Linear(self.hidden_size, self.hidden_size, bias=self.use_proj_bias) 

104 if len(self.weight_mode) > 0: 

105 self.weight_facet_decoder = nn.Linear(self.hidden_size * hidden_state_input_ratio, self.n_facet_effective) 

106 self.weight_global = nn.Parameter(torch.ones(self.n_facet_effective)) 

107 self.output_probs = True 

108 self.item_embedding = nn.Embedding(self.n_items, self.hidden_size, padding_idx=0) 

109 self.position_embedding = nn.Embedding(self.max_seq_length, self.hidden_size) 

110 self.trm_encoder = TransformerEncoder( 

111 n_layers=self.n_layers, 

112 n_heads=self.n_heads, 

113 hidden_size=self.hidden_size, 

114 inner_size=self.inner_size, 

115 hidden_dropout_prob=self.hidden_dropout_prob, 

116 attn_dropout_prob=self.attn_dropout_prob, 

117 hidden_act=self.hidden_act, 

118 layer_norm_eps=self.layer_norm_eps, 

119 ) 

120 

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

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

123 

124 if self.loss_type == "BPR": 

125 print("current softmax-cpr code does not support BPR loss") 

126 sys.exit(0) 

127 elif self.loss_type == "CE": 

128 self.loss_fct = nn.NLLLoss(reduction="none", ignore_index=0) # modified for mfs 

129 else: 

130 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!") 

131 

132 # parameters initialization 

133 self.apply(self._init_weights) 

134 

135 def get_facet_emb(self, input_emb, i): 

136 return self.project_arr[i](input_emb) 

137 

138 def _init_weights(self, module): 

139 """Initialize the weights""" 

140 if isinstance(module, (nn.Linear, nn.Embedding)): 

141 # Slightly different from the TF version which uses truncated_normal for initialization 

142 # cf https://github.com/pytorch/pytorch/pull/5617 

143 module.weight.data.normal_(mean=0.0, std=self.initializer_range) 

144 elif isinstance(module, nn.LayerNorm): 

145 module.bias.data.zero_() 

146 module.weight.data.fill_(1.0) 

147 if isinstance(module, nn.Linear) and module.bias is not None: 

148 module.bias.data.zero_() 

149 

150 def forward(self, item_seq, item_seq_len): 

151 position_ids = torch.arange(item_seq.size(1), dtype=torch.long, device=item_seq.device) 

152 position_ids = position_ids.unsqueeze(0).expand_as(item_seq) 

153 position_embedding = self.position_embedding(position_ids) 

154 

155 item_emb = self.item_embedding(item_seq) 

156 input_emb = item_emb + position_embedding 

157 input_emb = self.LayerNorm(input_emb) 

158 input_emb = self.dropout(input_emb) 

159 

160 extended_attention_mask = self.get_attention_mask(item_seq) 

161 

162 trm_output = self.trm_encoder(input_emb, extended_attention_mask, output_all_encoded_layers=True) 

163 return trm_output 

164 

165 def calculate_loss_prob(self, interaction, only_compute_prob=False): 

166 item_seq = interaction[self.ITEM_SEQ] 

167 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

168 all_hidden_states = self.forward(item_seq, item_seq_len) 

169 if self.loss_type != "CE": 

170 print("current softmax-cpr code does not support BPR or the losses other than cross entropy") 

171 sys.exit(0) 

172 else: # self.loss_type = 'CE' 

173 """mfs code starts""" 

174 device = all_hidden_states[0].device 

175 # check seq_len from hidden size 

176 

177 ## Multi-input hidden states: generate q_ct from hidden states 

178 # list of hidden state embeddings taken as input 

179 hidden_emb_arr = [] 

180 # h_facet_hidden -> H, n_face_window -> W, here 1 and 0 

181 for i in range(self.n_facet_hidden): 

182 hidden_states = all_hidden_states[-(i + 1)] # i-th hidden-state embedding from the top 

183 device = hidden_states.device 

184 hidden_emb_arr.append(hidden_states) 

185 for j in range(self.n_facet_window): 

186 ( 

187 bsz, 

188 seq_len, 

189 hidden_size, 

190 ) = hidden_states.size() # bsz -> , seq_len -> , hidden_size -> 768 in GPT-small? 

191 if j + 1 < hidden_states.size(1): 

192 shifted_hidden = torch.cat( 

193 ( 

194 torch.zeros((bsz, (j + 1), hidden_size), device=device), 

195 hidden_states[:, : -(j + 1), :], 

196 ), 

197 dim=1, 

198 ) 

199 else: 

200 shifted_hidden = torch.zeros((bsz, hidden_states.size(1), hidden_size), device=device) 

201 hidden_emb_arr.append(shifted_hidden) 

202 # hidden_emb_arr -> (W*H, bsz, seq_len, hidden_size) 

203 

204 # n_facet_MLP -> 1 

205 if self.n_facet_MLP > 0: 

206 stacked_hidden_emb_raw_arr = torch.cat(hidden_emb_arr, dim=-1) # (bsz, seq_len, W*H*hidden_size) 

207 # self.MLP_linear = nn.Linear( 

208 # config.hidden_size * (n_facet_hidden * (n_facet_window+1) ), # -> why +1? 

209 # config.hidden_size * n_facet_MLP 

210 # ) 

211 hidden_emb_MLP = self.MLP_linear(stacked_hidden_emb_raw_arr) # bsz, seq_len, hidden_size 

212 stacked_hidden_emb_arr_raw = torch.cat( 

213 [hidden_emb_arr[0], gelu(hidden_emb_MLP)], dim=-1 

214 ) # bsz, seq_len, 2*hidden_size 

215 else: 

216 stacked_hidden_emb_arr_raw = hidden_emb_arr[0] 

217 

218 # Only use the hidden state corresponding to the last item 

219 # The seq_len = 1 in the following code 

220 stacked_hidden_emb_arr = stacked_hidden_emb_arr_raw[:, -1, :].unsqueeze(dim=1) 

221 

222 # list of linear projects per facet 

223 projected_emb_arr = [] 

224 # list of final logits per facet 

225 facet_lm_logits_arr = [] 

226 

227 # logits for orig facets 

228 rereanker_candidate_token_ids_arr = [] 

229 for i in range(self.n_facet): 

230 # #linear projection 

231 projected_emb = self.get_facet_emb(stacked_hidden_emb_arr, i) # (bsz, seq_len, hidden_dim) 

232 projected_emb_arr.append(projected_emb) 

233 # logits for all tokens in vocab 

234 lm_logits = F.linear(projected_emb, self.item_embedding.weight, None) 

235 facet_lm_logits_arr.append(lm_logits) 

236 if i < self.n_facet_reranker and not self.candidates_from_previous_reranker: 

237 candidate_token_ids = [] 

238 for j in range(len(self.reranker_CAN_NUM)): 

239 _, candidate_token_ids_ = torch.topk(lm_logits, self.reranker_CAN_NUM[j]) 

240 candidate_token_ids.append(candidate_token_ids_) 

241 rereanker_candidate_token_ids_arr.append(candidate_token_ids) 

242 

243 for i in range(self.n_facet_reranker): 

244 for j in range(len(self.reranker_CAN_NUM)): 

245 projected_emb = self.get_facet_emb( 

246 stacked_hidden_emb_arr, 

247 self.n_facet + i * len(self.reranker_CAN_NUM) + j, 

248 ) # (bsz, seq_len, hidden_dim) 

249 projected_emb_arr.append(projected_emb) 

250 

251 for i in range(self.n_facet_context): 

252 projected_emb = self.get_facet_emb( 

253 stacked_hidden_emb_arr, 

254 self.n_facet + self.n_facet_reranker * len(self.reranker_CAN_NUM) + i, 

255 ) # (bsz, seq_len, hidden_dim) 

256 projected_emb_arr.append(projected_emb) 

257 

258 # to generate context-based embeddings for words in input 

259 for i in range(self.n_facet_emb): 

260 projected_emb = self.get_facet_emb( 

261 stacked_hidden_emb_arr_raw, 

262 self.n_facet + self.n_facet_context + self.n_facet_reranker * len(self.reranker_CAN_NUM) + i, 

263 ) # (bsz, seq_len, hidden_dim) 

264 projected_emb_arr.append(projected_emb) 

265 

266 for i in range(self.n_facet_reranker): 

267 bsz, seq_len, hidden_size = projected_emb_arr[i].size() 

268 for j in range(len(self.reranker_CAN_NUM)): 

269 if self.candidates_from_previous_reranker: 

270 _, candidate_token_ids = torch.topk( 

271 facet_lm_logits_arr[i], self.reranker_CAN_NUM[j] 

272 ) # (bsz, seq_len, topk) 

273 else: 

274 candidate_token_ids = rereanker_candidate_token_ids_arr[i][j] 

275 logit_hidden_reranker_topn = ( 

276 projected_emb_arr[self.n_facet + i * len(self.reranker_CAN_NUM) + j] 

277 .unsqueeze(dim=2) 

278 .expand(bsz, seq_len, self.reranker_CAN_NUM[j], hidden_size) 

279 * self.item_embedding.weight[candidate_token_ids, :] 

280 ).sum(dim=-1) # (bsz, seq_len, emb_size) x (bsz, seq_len, topk, emb_size) -> (bsz, seq_len, topk) 

281 if self.reranker_merging_mode == "add": 

282 facet_lm_logits_arr[i].scatter_add_( 

283 2, candidate_token_ids, logit_hidden_reranker_topn 

284 ) # (bsz, seq_len, vocab_size) <- (bsz, seq_len, topk) x (bsz, seq_len, topk) 

285 else: 

286 facet_lm_logits_arr[i].scatter_( 

287 2, candidate_token_ids, logit_hidden_reranker_topn 

288 ) # (bsz, seq_len, vocab_size) <- (bsz, seq_len, topk) x (bsz, seq_len, topk) 

289 

290 for i in range(self.n_facet_context): 

291 bsz, seq_len_1, hidden_size = projected_emb_arr[i].size() 

292 bsz, seq_len_2 = item_seq.size() 

293 logit_hidden_context = ( 

294 projected_emb_arr[self.n_facet + self.n_facet_reranker * len(self.reranker_CAN_NUM) + i] 

295 .unsqueeze(dim=2) 

296 .expand(-1, -1, seq_len_2, -1) 

297 * self.item_embedding.weight[item_seq, :].unsqueeze(dim=1).expand(-1, seq_len_1, -1, -1) 

298 ).sum(dim=-1) 

299 logit_hidden_pointer = 0 

300 if self.n_facet_emb == 2: # noqa: PLR2004 

301 logit_hidden_pointer = ( 

302 projected_emb_arr[-2][:, -1, :] 

303 .unsqueeze(dim=1) 

304 .unsqueeze(dim=1) 

305 .expand(-1, seq_len_1, seq_len_2, -1) 

306 * projected_emb_arr[-1].unsqueeze(dim=1).expand(-1, seq_len_1, -1, -1) 

307 ).sum(dim=-1) 

308 

309 item_seq_expand = item_seq.unsqueeze(dim=1).expand(-1, seq_len_1, -1) 

310 only_new_logits = torch.zeros_like(facet_lm_logits_arr[i]) 

311 if self.context_norm: 

312 only_new_logits.scatter_add_( 

313 dim=2, 

314 index=item_seq_expand, 

315 src=logit_hidden_context + logit_hidden_pointer, 

316 ) 

317 item_count = torch.zeros_like(only_new_logits) + 1e-15 

318 item_count.scatter_add_( 

319 dim=2, 

320 index=item_seq_expand, 

321 src=torch.ones_like(item_seq_expand).to(dtype=item_count.dtype), 

322 ) 

323 only_new_logits = only_new_logits / item_count 

324 else: 

325 only_new_logits.scatter_add_(dim=2, index=item_seq_expand, src=logit_hidden_context) 

326 item_count = torch.zeros_like(only_new_logits) + 1e-15 

327 item_count.scatter_add_( 

328 dim=2, 

329 index=item_seq_expand, 

330 src=torch.ones_like(item_seq_expand).to(dtype=item_count.dtype), 

331 ) 

332 only_new_logits = only_new_logits / item_count 

333 only_new_logits.scatter_add_(dim=2, index=item_seq_expand, src=logit_hidden_pointer) 

334 

335 if self.partition_merging_mode == "replace": 

336 facet_lm_logits_arr[i].scatter_( 

337 dim=2, 

338 index=item_seq_expand, 

339 src=torch.zeros_like(item_seq_expand).to(dtype=facet_lm_logits_arr[i].dtype), 

340 ) 

341 facet_lm_logits_arr[i] = facet_lm_logits_arr[i] + only_new_logits 

342 

343 weight = None 

344 if self.weight_mode == "dynamic": 

345 weight = self.weight_facet_decoder(stacked_hidden_emb_arr).softmax( 

346 dim=-1 

347 ) # hidden_dim*hidden_input_state_ration -> n_facet_effective 

348 elif self.weight_mode == "static": 

349 weight = self.weight_global.softmax(dim=-1) # torch.ones(n_facet_effective) 

350 elif self.weight_mode == "max_logits": 

351 stacked_facet_lm_logits = torch.stack(facet_lm_logits_arr, dim=0) 

352 facet_lm_logits_arr = [stacked_facet_lm_logits.amax(dim=0)] 

353 

354 prediction_prob = 0 

355 

356 for i in range(self.n_facet_effective): 

357 facet_lm_logits = facet_lm_logits_arr[i] 

358 if self.softmax_nonlinear == "sigsoftmax": #'None' here 

359 facet_lm_logits_sig = torch.exp(facet_lm_logits - facet_lm_logits.max(dim=-1, keepdim=True)[0]) * ( 

360 1e-20 + torch.sigmoid(facet_lm_logits) 

361 ) 

362 facet_lm_logits_softmax = facet_lm_logits_sig / facet_lm_logits_sig.sum(dim=-1, keepdim=True) 

363 elif self.softmax_nonlinear == "None": 

364 facet_lm_logits_softmax = facet_lm_logits.softmax(dim=-1) # softmax over final logits 

365 if self.weight_mode == "dynamic": 

366 prediction_prob += facet_lm_logits_softmax * weight[:, :, i].unsqueeze(-1) 

367 elif self.weight_mode == "static": 

368 prediction_prob += facet_lm_logits_softmax * weight[i] 

369 else: 

370 prediction_prob += facet_lm_logits_softmax / self.n_facet_effective # softmax over final logits/1 

371 if not only_compute_prob: 

372 inp = torch.log(prediction_prob.view(-1, self.n_items) + 1e-8) 

373 pos_items = interaction[self.POS_ITEM_ID] 

374 loss_raw = self.loss_fct(inp, pos_items.view(-1)) 

375 loss = loss_raw.mean() 

376 else: 

377 loss = None 

378 # return loss, prediction_prob.squeeze() 

379 return loss, prediction_prob.squeeze(dim=1) 

380 

381 def calculate_loss(self, interaction): 

382 loss, prediction_prob = self.calculate_loss_prob(interaction) 

383 return loss 

384 

385 def predict(self, interaction): 

386 print( 

387 "Current softmax cpr code does not support negative sampling in an efficient way just like RepeatNet.", 

388 file=sys.stderr, 

389 ) 

390 assert False # If you can accept slow running time, comment this line 

391 loss, prediction_prob = self.calculate_loss_prob(interaction, only_compute_prob=True) 

392 if self.post_remove_context: 

393 item_seq = interaction[self.ITEM_SEQ] 

394 prediction_prob.scatter_(1, item_seq, 0) 

395 test_item = interaction[self.ITEM_ID] 

396 prediction_prob = prediction_prob.unsqueeze(-1) 

397 # batch_size * num_items * 1 

398 scores = self.gather_indexes(prediction_prob, test_item).squeeze(-1) 

399 

400 return scores 

401 

402 def full_sort_predict(self, interaction): 

403 loss, prediction_prob = self.calculate_loss_prob(interaction) 

404 if self.post_remove_context: 

405 item_seq = interaction[self.ITEM_SEQ] 

406 prediction_prob.scatter_(1, item_seq, 0) 

407 return prediction_prob