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
« 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
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
10# UPDATE:
11# @Time : 2023/11/24
12# @Author : Haw-Shiuan Chang
13# @Email : ken77921@gmail.com
15r"""GRU4Rec + Softmax-CPR
16################################################
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
23""" # noqa: E501
25import math
26import sys
28import torch
29import torch.nn.functional as F
30from torch import nn
31from torch.nn.init import xavier_normal_, xavier_uniform_
33from hopwise.model.abstract_recommender import SequentialRecommender
36def gelu(x):
37 return 0.5 * x * (1 + torch.tanh(math.sqrt(2 / math.pi) * (x + 0.044715 * torch.pow(x, 3))))
40class GRU4RecCPR(SequentialRecommender):
41 r"""GRU4Rec is a model that incorporate RNN for recommendation.
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 """
49 def __init__(self, config, dataset):
50 super().__init__(config, dataset)
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"]
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
77 self.dense = nn.Linear(self.hidden_size, self.embedding_size)
78 out_size = self.embedding_size
80 self.n_embd = out_size
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
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
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 )
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))
117 self.c = 123
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)
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
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']!")
144 # parameters initialization
145 self.apply(self._init_weights)
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)
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
161 def get_facet_emb(self, input_emb, i):
162 return self.project_arr[i](input_emb)
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
176 """mfs code starts"""
177 device = all_hidden_states[0].device
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)
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]
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)
218 # list of linear projects per facet
219 projected_emb_arr = []
220 # list of final logits per facet
221 facet_lm_logits_arr = []
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)
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)
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)
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)
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)
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)
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)
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
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)]
362 prediction_prob = 0
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)
388 def calculate_loss(self, interaction):
389 loss, prediction_prob = self.calculate_loss_prob(interaction)
390 return loss
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
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