Coverage for hopwise/model/sequential_recommender/lightsans.py: 90%

84 statements  

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

1# @Time : 2021/05/01 

2# @Author : Xinyan Fan 

3# @Email : xinyan.fan@ruc.edu.cn 

4 

5"""LightSANs 

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

7Reference: 

8 Xin-Yan Fan et al. "Lighter and Better: Low-Rank Decomposed Self-Attention Networks for Next-Item Recommendation." in SIGIR 2021. 

9Reference: 

10 https://github.com/BELIEVEfxy/LightSANs 

11""" # noqa: E501 

12 

13import torch 

14from torch import nn 

15 

16from hopwise.model.abstract_recommender import SequentialRecommender 

17from hopwise.model.layers import LightTransformerEncoder 

18from hopwise.model.loss import BPRLoss 

19 

20 

21class LightSANs(SequentialRecommender): 

22 def __init__(self, config, dataset): 

23 super().__init__(config, dataset) 

24 

25 # load parameters info 

26 self.n_layers = config["n_layers"] 

27 self.n_heads = config["n_heads"] 

28 self.k_interests = config["k_interests"] 

29 self.hidden_size = config["hidden_size"] # same as embedding_size 

30 self.inner_size = config["inner_size"] # the dimensionality in feed-forward layer 

31 self.hidden_dropout_prob = config["hidden_dropout_prob"] 

32 self.attn_dropout_prob = config["attn_dropout_prob"] 

33 self.hidden_act = config["hidden_act"] 

34 self.layer_norm_eps = config["layer_norm_eps"] 

35 

36 self.initializer_range = config["initializer_range"] 

37 self.loss_type = config["loss_type"] 

38 

39 self.seq_len = self.max_seq_length 

40 # define layers and loss 

41 self.item_embedding = nn.Embedding(self.n_items, self.hidden_size, padding_idx=0) 

42 self.position_embedding = nn.Embedding(self.max_seq_length, self.hidden_size) 

43 self.trm_encoder = LightTransformerEncoder( 

44 n_layers=self.n_layers, 

45 n_heads=self.n_heads, 

46 k_interests=self.k_interests, 

47 hidden_size=self.hidden_size, 

48 seq_len=self.seq_len, 

49 inner_size=self.inner_size, 

50 hidden_dropout_prob=self.hidden_dropout_prob, 

51 attn_dropout_prob=self.attn_dropout_prob, 

52 hidden_act=self.hidden_act, 

53 layer_norm_eps=self.layer_norm_eps, 

54 ) 

55 

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

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

58 

59 if self.loss_type == "BPR": 

60 self.loss_fct = BPRLoss() 

61 elif self.loss_type == "CE": 

62 self.loss_fct = nn.CrossEntropyLoss() 

63 else: 

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

65 

66 # parameters initialization 

67 self.apply(self._init_weights) 

68 

69 def _init_weights(self, module): 

70 """Initialize the weights""" 

71 if isinstance(module, (nn.Linear, nn.Embedding)): 

72 module.weight.data.normal_(mean=0.0, std=self.initializer_range) 

73 elif isinstance(module, nn.LayerNorm): 

74 module.bias.data.zero_() 

75 module.weight.data.fill_(1.0) 

76 if isinstance(module, nn.Linear) and module.bias is not None: 

77 module.bias.data.zero_() 

78 

79 def embedding_layer(self, item_seq): 

80 position_ids = torch.arange(item_seq.size(1), dtype=torch.long, device=item_seq.device) 

81 position_embedding = self.position_embedding(position_ids) 

82 item_emb = self.item_embedding(item_seq) 

83 return item_emb, position_embedding 

84 

85 def forward(self, item_seq, item_seq_len): 

86 item_emb, position_embedding = self.embedding_layer(item_seq) 

87 item_emb = self.LayerNorm(item_emb) 

88 item_emb = self.dropout(item_emb) 

89 

90 trm_output = self.trm_encoder(item_emb, position_embedding, output_all_encoded_layers=True) 

91 output = trm_output[-1] 

92 output = self.gather_indexes(output, item_seq_len - 1) 

93 return output # [B H] 

94 

95 def calculate_loss(self, interaction): 

96 item_seq = interaction[self.ITEM_SEQ] 

97 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

98 seq_output = self.forward(item_seq, item_seq_len) 

99 pos_items = interaction[self.POS_ITEM_ID] 

100 if self.loss_type == "BPR": 

101 neg_items = interaction[self.NEG_ITEM_ID] 

102 pos_items_emb = self.item_embedding(pos_items) 

103 neg_items_emb = self.item_embedding(neg_items) 

104 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) # [B] 

105 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) # [B] 

106 loss = self.loss_fct(pos_score, neg_score) 

107 return loss 

108 else: # self.loss_type = 'CE' 

109 test_item_emb = self.item_embedding.weight 

110 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1)) 

111 loss = self.loss_fct(logits, pos_items) 

112 return loss 

113 

114 def predict(self, interaction): 

115 item_seq = interaction[self.ITEM_SEQ] 

116 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

117 test_item = interaction[self.ITEM_ID] 

118 

119 seq_output = self.forward(item_seq, item_seq_len) 

120 test_item_emb = self.item_embedding(test_item) 

121 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) # [B] 

122 return scores 

123 

124 def full_sort_predict(self, interaction): 

125 item_seq = interaction[self.ITEM_SEQ] 

126 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

127 seq_output = self.forward(item_seq, item_seq_len) 

128 test_items_emb = self.item_embedding.weight 

129 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) # [B n_items] 

130 return scores