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

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

2# @Author : Hui Wang 

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

4 

5"""SASRec 

6################################################ 

7 

8Reference: 

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

10 

11Reference: 

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

13 

14""" 

15 

16import torch 

17from torch import nn 

18 

19from hopwise.model.abstract_recommender import SequentialRecommender 

20from hopwise.model.layers import TransformerEncoder 

21from hopwise.model.loss import BPRLoss 

22 

23 

24class SASRec(SequentialRecommender): 

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

26 

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 """ 

32 

33 def __init__(self, config, dataset): 

34 super().__init__(config, dataset) 

35 

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"] 

45 

46 self.initializer_range = config["initializer_range"] 

47 self.loss_type = config["loss_type"] 

48 

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 ) 

62 

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

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

65 

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']!") 

72 

73 # parameters initialization 

74 self.apply(self._init_weights) 

75 

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_() 

87 

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) 

92 

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) 

97 

98 extended_attention_mask = self.get_attention_mask(item_seq) 

99 

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] 

104 

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 

123 

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 

132 

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