Coverage for hopwise/model/sequential_recommender/gru4reccpr.py: 81%

221 statements  

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

1# @Time : 2020/8/17 19:38 

2# @Author : Yujie Lu 

3# @Email : yujielu1998@gmail.com 

4 

5# UPDATE: 

6# @Time : 2020/8/19, 2020/10/2 

7# @Author : Yupeng Hou, Yujie Lu 

8# @Email : houyupeng@ruc.edu.cn, yujielu1998@gmail.com 

9 

10# UPDATE: 

11# @Time : 2023/11/24 

12# @Author : Haw-Shiuan Chang 

13# @Email : ken77921@gmail.com 

14 

15r"""GRU4Rec + Softmax-CPR 

16################################################ 

17 

18Reference: 

19 Yong Kiam Tan et al. "Improved Recurrent Neural Networks for Session-based Recommendations." in DLRS 2016. 

20 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 

21 

22 

23""" # noqa: E501 

24 

25import math 

26import sys 

27 

28import torch 

29import torch.nn.functional as F 

30from torch import nn 

31from torch.nn.init import xavier_normal_, xavier_uniform_ 

32 

33from hopwise.model.abstract_recommender import SequentialRecommender 

34 

35 

36def gelu(x): 

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

38 

39 

40class GRU4RecCPR(SequentialRecommender): 

41 r"""GRU4Rec is a model that incorporate RNN for recommendation. 

42 

43 Note: 

44 Regarding the innovation of this article,we can only achieve the data augmentation mentioned 

45 in the paper and directly output the embedding of the item, 

46 in order that the generation method we used is common to other sequential models. 

47 """ 

48 

49 def __init__(self, config, dataset): 

50 super().__init__(config, dataset) 

51 

52 # load parameters info 

53 self.hidden_size = config["hidden_size"] 

54 self.embedding_size = config["embedding_size"] 

55 self.loss_type = config["loss_type"] 

56 self.num_layers = config["num_layers"] 

57 self.dropout_prob = config["dropout_prob"] 

58 

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

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

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

62 self.n_facet_hidden = min( 

63 config["n_facet_hidden"], config["num_layers"] 

64 ) # config['n_facet_hidden'] #added for mfs 

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

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

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

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

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

70 assert self.n_facet_window <= 0 

71 self.n_facet_window = -self.n_facet_window 

72 self.n_facet_MLP = -self.n_facet_MLP 

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

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

75 self.only_compute_loss = True # added for mfs 

76 

77 self.dense = nn.Linear(self.hidden_size, self.embedding_size) 

78 out_size = self.embedding_size 

79 

80 self.n_embd = out_size 

81 

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

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

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

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

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

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

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

89 assert self.use_proj_bias is not None 

90 self.candidates_from_previous_reranker = True 

91 if self.weight_mode == "max_logits": 

92 self.n_facet_effective = 1 

93 else: 

94 self.n_facet_effective = self.n_facet 

95 

96 assert ( 

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

98 == self.n_facet_all 

99 ) 

100 assert self.n_facet_emb in (0, 2) # noqa: PLR2004 

101 

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

103 self.MLP_linear = nn.Linear( 

104 self.n_embd * (self.n_facet_hidden * (self.n_facet_window + 1)), 

105 self.n_embd * self.n_facet_MLP, 

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

107 total_lin_dim = self.n_embd * hidden_state_input_ratio 

108 self.project_arr = nn.ModuleList( 

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

110 ) 

111 

112 self.project_emb = nn.Linear(self.n_embd, self.n_embd, bias=self.use_proj_bias) 

113 if len(self.weight_mode) > 0: 

114 self.weight_facet_decoder = nn.Linear(self.n_embd * hidden_state_input_ratio, self.n_facet_effective) 

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

116 

117 self.c = 123 

118 

119 # define layers and loss 

120 self.emb_dropout = nn.Dropout(self.dropout_prob) 

121 self.gru_layers = nn.GRU( 

122 input_size=self.embedding_size, 

123 hidden_size=self.hidden_size, 

124 num_layers=self.num_layers, 

125 bias=False, 

126 batch_first=True, 

127 ) 

128 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0) 

129 

130 if self.use_out_emb: 

131 self.out_item_embedding = nn.Linear(out_size, self.n_items, bias=False) 

132 else: 

133 self.out_item_embedding = self.item_embedding 

134 self.out_item_embedding.bias = None 

135 

136 if self.loss_type == "BPR": 

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

138 sys.exit(0) 

139 elif self.loss_type == "CE": 

140 self.loss_fct = nn.CrossEntropyLoss() 

141 else: 

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

143 

144 # parameters initialization 

145 self.apply(self._init_weights) 

146 

147 def _init_weights(self, module): 

148 if isinstance(module, nn.Embedding): 

149 xavier_normal_(module.weight) 

150 elif isinstance(module, nn.GRU): 

151 xavier_uniform_(module.weight_hh_l0) 

152 xavier_uniform_(module.weight_ih_l0) 

153 

154 def forward(self, item_seq, item_seq_len): 

155 item_seq_emb = self.item_embedding(item_seq) 

156 item_seq_emb_dropout = self.emb_dropout(item_seq_emb) 

157 gru_output, _ = self.gru_layers(item_seq_emb_dropout) 

158 gru_output = self.dense(gru_output) 

159 return gru_output 

160 

161 def get_facet_emb(self, input_emb, i): 

162 return self.project_arr[i](input_emb) 

163 

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

165 item_seq = interaction[self.ITEM_SEQ] 

166 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

167 last_layer_hs = self.forward(item_seq, item_seq_len) 

168 all_hidden_states = [last_layer_hs] 

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 test_item_emb = self.out_item_embedding.weight 

174 test_item_bias = self.out_item_embedding.bias 

175 

176 """mfs code starts""" 

177 device = all_hidden_states[0].device 

178 

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 # print('all_hidden_states length is {}. i is {}'.format(len(all_hidden_states), i)) 

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

184 device = hidden_states.device 

185 hidden_emb_arr.append(hidden_states) 

186 for j in range(self.n_facet_window): 

187 ( 

188 bsz, 

189 seq_len, 

190 hidden_size, 

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

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

193 shifted_hidden = torch.cat( 

194 ( 

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

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

197 ), 

198 dim=1, 

199 ) 

200 else: 

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

202 hidden_emb_arr.append(shifted_hidden) 

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

204 

205 # n_facet_MLP -> 1 

206 if self.n_facet_MLP > 0: 

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

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

209 stacked_hidden_emb_arr_raw = torch.cat( 

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

211 ) # bsz, seq_len, 2*hidden_size 

212 else: 

213 stacked_hidden_emb_arr_raw = hidden_emb_arr[0] 

214 

215 # Only use the hidden state corresponding to the last word 

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

217 

218 # list of linear projects per facet 

219 projected_emb_arr = [] 

220 # list of final logits per facet 

221 facet_lm_logits_arr = [] 

222 

223 rereanker_candidate_token_ids_arr = [] 

224 for i in range(self.n_facet): 

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

226 projected_emb_arr.append(projected_emb) 

227 lm_logits = F.linear(projected_emb, test_item_emb, test_item_bias) 

228 facet_lm_logits_arr.append(lm_logits) 

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

230 candidate_token_ids = [] 

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

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

233 candidate_token_ids.append(candidate_token_ids_) 

234 rereanker_candidate_token_ids_arr.append(candidate_token_ids) 

235 

236 for i in range(self.n_facet_reranker): 

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

238 projected_emb = self.get_facet_emb( 

239 stacked_hidden_emb_arr, 

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

241 ) # (bsz, seq_len, hidden_dim) 

242 projected_emb_arr.append(projected_emb) 

243 

244 for i in range(self.n_facet_context): 

245 projected_emb = self.get_facet_emb( 

246 stacked_hidden_emb_arr, 

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

248 ) # (bsz, seq_len, hidden_dim) 

249 projected_emb_arr.append(projected_emb) 

250 

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

252 for i in range(self.n_facet_emb): 

253 projected_emb = self.get_facet_emb( 

254 stacked_hidden_emb_arr_raw, 

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

256 ) # (bsz, seq_len, hidden_dim) 

257 projected_emb_arr.append(projected_emb) 

258 

259 for i in range(self.n_facet_reranker): 

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

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

262 if self.candidates_from_previous_reranker: 

263 _, candidate_token_ids = torch.topk( 

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

265 ) # (bsz, seq_len, topk) 

266 else: 

267 candidate_token_ids = rereanker_candidate_token_ids_arr[i][j] 

268 logit_hidden_reranker_topn = ( 

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

270 .unsqueeze(dim=2) 

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

272 * test_item_emb[candidate_token_ids, :] 

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

274 if test_item_bias is not None: 

275 logit_hidden_reranker_topn += test_item_bias[candidate_token_ids] 

276 if self.reranker_merging_mode == "add": 

277 # print("inside reranker") 

278 facet_lm_logits_arr[i].scatter_add_( 

279 2, candidate_token_ids, logit_hidden_reranker_topn 

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

281 else: 

282 facet_lm_logits_arr[i].scatter_( 

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 

286 for i in range(self.n_facet_context): 

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

288 bsz, seq_len_2 = item_seq.size() 

289 logit_hidden_context = ( 

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

291 .unsqueeze(dim=2) 

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

293 * test_item_emb[item_seq, :].unsqueeze(dim=1).expand(-1, seq_len_1, -1, -1) 

294 ).sum(dim=-1) 

295 if test_item_bias is not None: 

296 logit_hidden_context += test_item_bias[item_seq].unsqueeze(dim=1).expand(-1, seq_len_1, -1) 

297 logit_hidden_pointer = 0 

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

299 logit_hidden_pointer = ( 

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

301 .unsqueeze(dim=1) 

302 .unsqueeze(dim=1) 

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

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

305 ).sum(dim=-1) 

306 

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

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

309 if self.context_norm: 

310 only_new_logits.scatter_add_( 

311 dim=2, 

312 index=item_seq_expand, 

313 src=logit_hidden_context + logit_hidden_pointer, 

314 ) 

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

316 item_count.scatter_add_( 

317 dim=2, 

318 index=item_seq_expand, 

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

320 ) 

321 only_new_logits = only_new_logits / item_count 

322 else: 

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

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

325 item_count.scatter_add_( 

326 dim=2, 

327 index=item_seq_expand, 

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

329 ) 

330 only_new_logits = only_new_logits / item_count 

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

332 

333 if self.partition_merging_mode == "replace": 

334 facet_lm_logits_arr[i].scatter_( 

335 dim=2, 

336 index=item_seq_expand, 

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

338 ) 

339 facet_lm_logits_arr[i] = facet_lm_logits_arr[i] + only_new_logits 

340 elif self.partition_merging_mode == "add": 

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

342 elif self.partition_merging_mode == "half": 

343 item_in_context = torch.ones_like(only_new_logits) 

344 item_in_context.scatter_( 

345 dim=2, 

346 index=item_seq_expand, 

347 src=2 * torch.ones_like(item_seq_expand).to(dtype=item_count.dtype), 

348 ) 

349 facet_lm_logits_arr[i] = facet_lm_logits_arr[i] / item_in_context + only_new_logits 

350 

351 weight = None 

352 if self.weight_mode == "dynamic": 

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

354 dim=-1 

355 ) # hidden_dim*hidden_input_state_ration -> n_facet_effective 

356 elif self.weight_mode == "static": 

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

358 elif self.weight_mode == "max_logits": 

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

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

361 

362 prediction_prob = 0 

363 

364 for i in range(self.n_facet_effective): 

365 facet_lm_logits = facet_lm_logits_arr[i] 

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

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

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

369 ) 

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

371 elif self.softmax_nonlinear == "None": 

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

373 if self.weight_mode == "dynamic": 

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

375 elif self.weight_mode == "static": 

376 prediction_prob += facet_lm_logits_softmax * weight[i] 

377 else: 

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

379 if not only_compute_prob: 

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

381 pos_items = interaction[self.POS_ITEM_ID] 

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

383 loss = loss_raw.mean() 

384 else: 

385 loss = None 

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

387 

388 def calculate_loss(self, interaction): 

389 loss, prediction_prob = self.calculate_loss_prob(interaction) 

390 return loss 

391 

392 def predict(self, interaction): 

393 print( 

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

395 file=sys.stderr, 

396 ) 

397 assert False # If you can accept slow running time, uncomment this line. 

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

399 if self.post_remove_context: 

400 item_seq = interaction[self.ITEM_SEQ] 

401 prediction_prob.scatter_(1, item_seq, 0) 

402 test_item = interaction[self.ITEM_ID] 

403 prediction_prob = prediction_prob.unsqueeze(-1) 

404 # batch_size * num_items * 1 

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

406 return scores 

407 

408 def full_sort_predict(self, interaction): 

409 loss, prediction_prob = self.calculate_loss_prob(interaction) 

410 item_seq = interaction[self.ITEM_SEQ] 

411 if self.post_remove_context: 

412 prediction_prob.scatter_(1, item_seq, 0) 

413 return prediction_prob