Coverage for hopwise/model/sequential_recommender/shan.py: 91%

107 statements  

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

1# @Time : 2020/11/20 22:33 

2# @Author : Shao Weiqi 

3# @Reviewer : Lin Kun 

4# @Email : shaoweiqi@ruc.edu.cn 

5 

6r"""SHAN 

7################################################ 

8 

9Reference: 

10 Ying, H et al. "Sequential Recommender System based on Hierarchical Attention Network."in IJCAI 2018 

11 

12 

13""" 

14 

15import numpy as np 

16import torch 

17from torch import nn 

18from torch.nn.init import normal_, uniform_ 

19 

20from hopwise.model.abstract_recommender import SequentialRecommender 

21from hopwise.model.loss import BPRLoss 

22 

23 

24class SHAN(SequentialRecommender): 

25 r"""SHAN exploit the Hierarchical Attention Network to get the long-short term preference 

26 first get the long term purpose and then fuse the long-term with recent items to get long-short term purpose 

27 

28 """ 

29 

30 def __init__(self, config, dataset): 

31 super().__init__(config, dataset) 

32 

33 # load the dataset information 

34 self.n_users = dataset.num(self.USER_ID) 

35 self.device = config["device"] 

36 self.INVERSE_ITEM_SEQ = config["INVERSE_ITEM_SEQ"] 

37 

38 # load the parameter information 

39 self.embedding_size = config["embedding_size"] 

40 self.short_item_length = config["short_item_length"] # the length of the short session items 

41 assert self.short_item_length <= self.max_seq_length, "short_item_length can't longer than the max_seq_length" 

42 self.reg_weight = config["reg_weight"] 

43 

44 # define layers and loss 

45 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0) 

46 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size) 

47 

48 self.long_w = nn.Linear(self.embedding_size, self.embedding_size) 

49 self.long_b = nn.Parameter( 

50 uniform_( 

51 tensor=torch.zeros(self.embedding_size), 

52 a=-np.sqrt(3 / self.embedding_size), 

53 b=np.sqrt(3 / self.embedding_size), 

54 ), 

55 requires_grad=True, 

56 ) 

57 self.long_short_w = nn.Linear(self.embedding_size, self.embedding_size) 

58 self.long_short_b = nn.Parameter( 

59 uniform_( 

60 tensor=torch.zeros(self.embedding_size), 

61 a=-np.sqrt(3 / self.embedding_size), 

62 b=np.sqrt(3 / self.embedding_size), 

63 ), 

64 requires_grad=True, 

65 ) 

66 

67 self.relu = nn.ReLU() 

68 

69 self.loss_type = config["loss_type"] 

70 if self.loss_type == "BPR": 

71 self.loss_fct = BPRLoss() 

72 elif self.loss_type == "CE": 

73 self.loss_fct = nn.CrossEntropyLoss() 

74 else: 

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

76 

77 # init the parameter of the model 

78 self.apply(self.init_weights) 

79 

80 def reg_loss(self, user_embedding, item_embedding): 

81 reg_1, reg_2 = self.reg_weight 

82 loss_1 = reg_1 * torch.norm(self.long_w.weight, p=2) + reg_1 * torch.norm(self.long_short_w.weight, p=2) 

83 loss_2 = reg_2 * torch.norm(user_embedding, p=2) + reg_2 * torch.norm(item_embedding, p=2) 

84 

85 return loss_1 + loss_2 

86 

87 def init_weights(self, module): 

88 if isinstance(module, nn.Embedding): 

89 normal_(module.weight.data, 0.0, 0.01) 

90 elif isinstance(module, nn.Linear): 

91 uniform_( 

92 module.weight.data, 

93 -np.sqrt(3 / self.embedding_size), 

94 np.sqrt(3 / self.embedding_size), 

95 ) 

96 elif isinstance(module, nn.Parameter): 

97 uniform_( 

98 module.data, 

99 -np.sqrt(3 / self.embedding_size), 

100 np.sqrt(3 / self.embedding_size), 

101 ) 

102 print(module.data) 

103 

104 def forward(self, seq_item, user): 

105 seq_item_embedding = self.item_embedding(seq_item) 

106 user_embedding = self.user_embedding(user) 

107 

108 # get the mask 

109 mask = seq_item.data.eq(0) 

110 long_term_attention_based_pooling_layer = self.long_term_attention_based_pooling_layer( 

111 seq_item_embedding, user_embedding, mask 

112 ) 

113 # batch_size * 1 * embedding_size 

114 

115 short_item_embedding = seq_item_embedding[:, -self.short_item_length :, :] 

116 mask_long_short = mask[:, -self.short_item_length :] 

117 batch_size = mask_long_short.size(0) 

118 x = torch.zeros(size=(batch_size, 1)).eq(1).to(self.device) 

119 mask_long_short = torch.cat([x, mask_long_short], dim=1) 

120 # batch_size * short_item_length * embedding_size 

121 long_short_item_embedding = torch.cat([long_term_attention_based_pooling_layer, short_item_embedding], dim=1) 

122 # batch_size * 1_plus_short_item_length * embedding_size 

123 

124 long_short_item_embedding = self.long_and_short_term_attention_based_pooling_layer( 

125 long_short_item_embedding, user_embedding, mask_long_short 

126 ) 

127 # batch_size * embedding_size 

128 

129 return long_short_item_embedding 

130 

131 def calculate_loss(self, interaction): 

132 inverse_seq_item = interaction[self.INVERSE_ITEM_SEQ] 

133 user = interaction[self.USER_ID] 

134 user_embedding = self.user_embedding(user) 

135 seq_output = self.forward(inverse_seq_item, user) 

136 pos_items = interaction[self.POS_ITEM_ID] 

137 pos_items_emb = self.item_embedding(pos_items) 

138 if self.loss_type == "BPR": 

139 neg_items = interaction[self.NEG_ITEM_ID] 

140 neg_items_emb = self.item_embedding(neg_items) 

141 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) 

142 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) 

143 loss = self.loss_fct(pos_score, neg_score) 

144 return loss + self.reg_loss(user_embedding, pos_items_emb) 

145 else: # self.loss_type = 'CE' 

146 test_item_emb = self.item_embedding.weight 

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

148 loss = self.loss_fct(logits, pos_items) 

149 return loss + self.reg_loss(user_embedding, pos_items_emb) 

150 

151 def predict(self, interaction): 

152 inverse_item_seq = interaction[self.INVERSE_ITEM_SEQ] 

153 test_item = interaction[self.ITEM_ID] 

154 user = interaction[self.USER_ID] 

155 seq_output = self.forward(inverse_item_seq, user) 

156 test_item_emb = self.item_embedding(test_item) 

157 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) 

158 return scores 

159 

160 def full_sort_predict(self, interaction): 

161 inverse_item_seq = interaction[self.ITEM_SEQ] 

162 user = interaction[self.USER_ID] 

163 seq_output = self.forward(inverse_item_seq, user) 

164 test_items_emb = self.item_embedding.weight 

165 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) 

166 return scores 

167 

168 def long_and_short_term_attention_based_pooling_layer(self, long_short_item_embedding, user_embedding, mask=None): 

169 """Fusing the long term purpose with the short-term preference""" 

170 long_short_item_embedding_value = long_short_item_embedding 

171 

172 long_short_item_embedding = self.relu(self.long_short_w(long_short_item_embedding) + self.long_short_b) 

173 long_short_item_embedding = torch.matmul(long_short_item_embedding, user_embedding.unsqueeze(2)).squeeze(-1) 

174 # batch_size * seq_len 

175 if mask is not None: 

176 long_short_item_embedding.masked_fill_(mask, -1e9) 

177 long_short_item_embedding = nn.Softmax(dim=-1)(long_short_item_embedding) 

178 long_short_item_embedding = torch.mul( 

179 long_short_item_embedding_value, long_short_item_embedding.unsqueeze(2) 

180 ).sum(dim=1) 

181 

182 return long_short_item_embedding 

183 

184 def long_term_attention_based_pooling_layer(self, seq_item_embedding, user_embedding, mask=None): 

185 """Get the long term purpose of user""" 

186 seq_item_embedding_value = seq_item_embedding 

187 

188 seq_item_embedding = self.relu(self.long_w(seq_item_embedding) + self.long_b) 

189 user_item_embedding = torch.matmul(seq_item_embedding, user_embedding.unsqueeze(2)).squeeze(-1) 

190 # batch_size * seq_len 

191 if mask is not None: 

192 user_item_embedding.masked_fill_(mask, -1e9) 

193 user_item_embedding = nn.Softmax(dim=1)(user_item_embedding) 

194 user_item_embedding = torch.mul(seq_item_embedding_value, user_item_embedding.unsqueeze(2)).sum( 

195 dim=1, keepdim=True 

196 ) 

197 # batch_size * 1 * embedding_size 

198 

199 return user_item_embedding