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
« 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
5# UPDATE:
6# @Time : 2023/11/24
7# @Author : Haw-Shiuan Chang
8# @Email : ken77921@gmail.com
10"""SASRec + Softmax-CPR
11################################################
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
17Reference:
18 https://github.com/kang205/SASRec
19 https://arxiv.org/pdf/2310.14079.pdf
21""" # noqa: E501
23# from hopwise.model.loss import BPRLoss
24import math
25import sys
27import torch
28import torch.nn.functional as F
29from torch import nn
31from hopwise.model.abstract_recommender import SequentialRecommender
32from hopwise.model.layers import TransformerEncoder
35def gelu(x):
36 return 0.5 * x * (1 + torch.tanh(math.sqrt(2 / math.pi) * (x + 0.044715 * torch.pow(x, 3))))
39class SASRecCPR(SequentialRecommender):
40 r"""SASRec is the first sequential recommender based on self-attentive mechanism.
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 """
48 def __init__(self, config, dataset):
49 super().__init__(config, dataset)
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
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 )
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 )
121 self.LayerNorm = nn.LayerNorm(self.hidden_size, eps=self.layer_norm_eps)
122 self.dropout = nn.Dropout(self.hidden_dropout_prob)
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']!")
132 # parameters initialization
133 self.apply(self._init_weights)
135 def get_facet_emb(self, input_emb, i):
136 return self.project_arr[i](input_emb)
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_()
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)
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)
160 extended_attention_mask = self.get_attention_mask(item_seq)
162 trm_output = self.trm_encoder(input_emb, extended_attention_mask, output_all_encoded_layers=True)
163 return trm_output
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
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)
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]
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)
222 # list of linear projects per facet
223 projected_emb_arr = []
224 # list of final logits per facet
225 facet_lm_logits_arr = []
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)
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)
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)
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)
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)
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)
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)
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
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)]
354 prediction_prob = 0
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)
381 def calculate_loss(self, interaction):
382 loss, prediction_prob = self.calculate_loss_prob(interaction)
383 return loss
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)
400 return scores
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