Coverage for hopwise/model/sequential_recommender/sasrec.py: 90%
82 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"""SASRec
6################################################
8Reference:
9 Wang-Cheng Kang et al. "Self-Attentive Sequential Recommendation." in ICDM 2018.
11Reference:
12 https://github.com/kang205/SASRec
14"""
16import torch
17from torch import nn
19from hopwise.model.abstract_recommender import SequentialRecommender
20from hopwise.model.layers import TransformerEncoder
21from hopwise.model.loss import BPRLoss
24class SASRec(SequentialRecommender):
25 r"""SASRec is the first sequential recommender based on self-attentive mechanism.
27 Note:
28 In the author's implementation, the Point-Wise Feed-Forward Network (PFFN) is implemented
29 by CNN with 1x1 kernel. In this implementation, we follows the original BERT implementation
30 using Fully Connected Layer to implement the PFFN.
31 """
33 def __init__(self, config, dataset):
34 super().__init__(config, dataset)
36 # load parameters info
37 self.n_layers = config["n_layers"]
38 self.n_heads = config["n_heads"]
39 self.hidden_size = config["hidden_size"] # same as embedding_size
40 self.inner_size = config["inner_size"] # the dimensionality in feed-forward layer
41 self.hidden_dropout_prob = config["hidden_dropout_prob"]
42 self.attn_dropout_prob = config["attn_dropout_prob"]
43 self.hidden_act = config["hidden_act"]
44 self.layer_norm_eps = config["layer_norm_eps"]
46 self.initializer_range = config["initializer_range"]
47 self.loss_type = config["loss_type"]
49 # define layers and loss
50 self.item_embedding = nn.Embedding(self.n_items, self.hidden_size, padding_idx=0)
51 self.position_embedding = nn.Embedding(self.max_seq_length, self.hidden_size)
52 self.trm_encoder = TransformerEncoder(
53 n_layers=self.n_layers,
54 n_heads=self.n_heads,
55 hidden_size=self.hidden_size,
56 inner_size=self.inner_size,
57 hidden_dropout_prob=self.hidden_dropout_prob,
58 attn_dropout_prob=self.attn_dropout_prob,
59 hidden_act=self.hidden_act,
60 layer_norm_eps=self.layer_norm_eps,
61 )
63 self.LayerNorm = nn.LayerNorm(self.hidden_size, eps=self.layer_norm_eps)
64 self.dropout = nn.Dropout(self.hidden_dropout_prob)
66 if self.loss_type == "BPR":
67 self.loss_fct = BPRLoss()
68 elif self.loss_type == "CE":
69 self.loss_fct = nn.CrossEntropyLoss()
70 else:
71 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!")
73 # parameters initialization
74 self.apply(self._init_weights)
76 def _init_weights(self, module):
77 """Initialize the weights"""
78 if isinstance(module, (nn.Linear, nn.Embedding)):
79 # Slightly different from the TF version which uses truncated_normal for initialization
80 # cf https://github.com/pytorch/pytorch/pull/5617
81 module.weight.data.normal_(mean=0.0, std=self.initializer_range)
82 elif isinstance(module, nn.LayerNorm):
83 module.bias.data.zero_()
84 module.weight.data.fill_(1.0)
85 if isinstance(module, nn.Linear) and module.bias is not None:
86 module.bias.data.zero_()
88 def forward(self, item_seq, item_seq_len):
89 position_ids = torch.arange(item_seq.size(1), dtype=torch.long, device=item_seq.device)
90 position_ids = position_ids.unsqueeze(0).expand_as(item_seq)
91 position_embedding = self.position_embedding(position_ids)
93 item_emb = self.item_embedding(item_seq)
94 input_emb = item_emb + position_embedding
95 input_emb = self.LayerNorm(input_emb)
96 input_emb = self.dropout(input_emb)
98 extended_attention_mask = self.get_attention_mask(item_seq)
100 trm_output = self.trm_encoder(input_emb, extended_attention_mask, output_all_encoded_layers=True)
101 output = trm_output[-1]
102 output = self.gather_indexes(output, item_seq_len - 1)
103 return output # [B H]
105 def calculate_loss(self, interaction):
106 item_seq = interaction[self.ITEM_SEQ]
107 item_seq_len = interaction[self.ITEM_SEQ_LEN]
108 seq_output = self.forward(item_seq, item_seq_len)
109 pos_items = interaction[self.POS_ITEM_ID]
110 if self.loss_type == "BPR":
111 neg_items = interaction[self.NEG_ITEM_ID]
112 pos_items_emb = self.item_embedding(pos_items)
113 neg_items_emb = self.item_embedding(neg_items)
114 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) # [B]
115 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) # [B]
116 loss = self.loss_fct(pos_score, neg_score)
117 return loss
118 else: # self.loss_type = 'CE'
119 test_item_emb = self.item_embedding.weight
120 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1))
121 loss = self.loss_fct(logits, pos_items)
122 return loss
124 def predict(self, interaction):
125 item_seq = interaction[self.ITEM_SEQ]
126 item_seq_len = interaction[self.ITEM_SEQ_LEN]
127 test_item = interaction[self.ITEM_ID]
128 seq_output = self.forward(item_seq, item_seq_len)
129 test_item_emb = self.item_embedding(test_item)
130 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) # [B]
131 return scores
133 def full_sort_predict(self, interaction):
134 item_seq = interaction[self.ITEM_SEQ]
135 item_seq_len = interaction[self.ITEM_SEQ_LEN]
136 seq_output = self.forward(item_seq, item_seq_len)
137 test_items_emb = self.item_embedding.weight
138 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) # [B n_items]
139 return scores