Coverage for hopwise/model/sequential_recommender/sine.py: 94%

129 statements  

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

1# @Time : 2021/11/23 11:10 

2# @Author : Jingqi Gao 

3# @Email : jgaoaz@connect.ust.hk 

4 

5r"""SINE 

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

7 

8Reference: 

9 Qiaoyu Tan et al. "Sparse-Interest Network for Sequential Recommendation." in WSDM 2021. 

10 

11""" 

12 

13import numpy as np 

14import torch 

15import torch.nn.functional as F 

16from torch import nn 

17from torch.nn.init import xavier_normal_ 

18 

19from hopwise.model.abstract_recommender import SequentialRecommender 

20from hopwise.model.loss import BPRLoss 

21from hopwise.utils import InputType 

22 

23torch.autograd.set_detect_anomaly(True) 

24 

25 

26class SINE(SequentialRecommender): 

27 input_type = InputType.PAIRWISE 

28 

29 def __init__(self, config, dataset): 

30 super().__init__(config, dataset) 

31 

32 # load dataset info 

33 self.n_users = dataset.user_num 

34 self.n_items = dataset.item_num 

35 

36 # load parameters info 

37 self.device = config["device"] 

38 self.embedding_size = config["embedding_size"] 

39 self.loss_type = config["loss_type"] 

40 self.layer_norm_eps = config["layer_norm_eps"] 

41 

42 if self.loss_type == "BPR": 

43 self.loss_fct = BPRLoss() 

44 elif self.loss_type == "CE": 

45 self.loss_fct = nn.CrossEntropyLoss() 

46 elif self.loss_type == "NLL": 

47 self.loss_fct = nn.NLLLoss() 

48 else: 

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

50 

51 self.D = config["embedding_size"] 

52 self.L = config["prototype_size"] # 500 for movie-len dataset 

53 self.k = config["interest_size"] # 4 for movie-len dataset 

54 self.tau = config["tau_ratio"] # 0.1 in paper 

55 self.reg_loss_ratio = config["reg_loss_ratio"] # 0.1 in paper 

56 

57 self.initializer_range = 0.01 

58 

59 self.w1 = self._init_weight((self.D, self.D)) 

60 self.w2 = self._init_weight(self.D) 

61 self.w3 = self._init_weight((self.D, self.D)) 

62 self.w4 = self._init_weight(self.D) 

63 

64 self.C = nn.Embedding(self.L, self.D) 

65 

66 self.w_k_1 = self._init_weight((self.k, self.D, self.D)) 

67 self.w_k_2 = self._init_weight((self.k, self.D)) 

68 self.item_embedding = nn.Embedding(self.n_items, self.D, padding_idx=0) 

69 self.ln2 = nn.LayerNorm(self.embedding_size, eps=self.layer_norm_eps) 

70 self.ln4 = nn.LayerNorm(self.embedding_size, eps=self.layer_norm_eps) 

71 

72 # parameters initialization 

73 self.apply(self._init_weights) 

74 

75 def _init_weight(self, shape): 

76 mat = torch.FloatTensor(np.random.normal(0, self.initializer_range, shape)) 

77 return nn.Parameter(mat, requires_grad=True) 

78 

79 def _init_weights(self, module): 

80 if isinstance(module, nn.Embedding): 

81 xavier_normal_(module.weight) 

82 elif isinstance(module, nn.LayerNorm): 

83 module.bias.data.zero_() 

84 module.weight.data.fill_(1.0) 

85 

86 def calculate_loss(self, interaction): 

87 item_seq = interaction[self.ITEM_SEQ] 

88 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

89 seq_output = self.forward(item_seq, item_seq_len) 

90 pos_items = interaction[self.POS_ITEM_ID] 

91 

92 if self.loss_type == "BPR": 

93 neg_items = interaction[self.NEG_ITEM_ID] 

94 pos_items_emb = self.item_embedding(pos_items) 

95 neg_items_emb = self.item_embedding(neg_items) 

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

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

98 loss = self.loss_fct(pos_score, neg_score) 

99 return loss 

100 elif self.loss_type == "CE": 

101 test_item_emb = self.item_embedding.weight 

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

103 loss = self.loss_fct(logits, pos_items) 

104 return loss 

105 else: 

106 test_item_emb = self.item_embedding.weight 

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

108 logits = F.log_softmax(logits, dim=1) 

109 loss = self.loss_fct(logits, pos_items) 

110 return loss + self.calculate_reg_loss() * self.reg_loss_ratio 

111 

112 def calculate_reg_loss(self): 

113 C_mean = torch.mean(self.C.weight, dim=1, keepdim=True) 

114 C_reg = self.C.weight - C_mean 

115 C_reg = C_reg.matmul(C_reg.T) / self.D 

116 return (torch.norm(C_reg) ** 2 - torch.norm(torch.diag(C_reg)) ** 2) / 2 

117 

118 def predict(self, interaction): 

119 item_seq = interaction[self.ITEM_SEQ] 

120 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

121 test_item = interaction[self.ITEM_ID] 

122 seq_output = self.forward(item_seq, item_seq_len) 

123 test_item_emb = self.item_embedding(test_item) 

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

125 return scores 

126 

127 def forward(self, item_seq, item_seq_len): 

128 x_u = self.item_embedding(item_seq).to(self.device) # [B, N, D] 

129 

130 # concept activation 

131 # sort by inner product 

132 x = torch.matmul(x_u, self.w1) 

133 x = torch.tanh(x) 

134 x = torch.matmul(x, self.w2) 

135 a = F.softmax(x, dim=1) 

136 z_u = torch.matmul(a.unsqueeze(2).transpose(1, 2), x_u).transpose(1, 2) 

137 s_u = torch.matmul(self.C.weight, z_u) 

138 s_u = s_u.squeeze(2) 

139 idx = s_u.argsort(1)[:, -self.k :] 

140 s_u_idx = s_u.sort(1)[0][:, -self.k :] 

141 c_u = self.C(idx) 

142 sigs = torch.sigmoid(s_u_idx.unsqueeze(2).repeat(1, 1, self.embedding_size)) 

143 C_u = c_u.mul(sigs) 

144 

145 # intention assignment 

146 # use matrix multiplication instead of cos() 

147 w3_x_u_norm = F.normalize(x_u.matmul(self.w3), p=2, dim=2) 

148 C_u_norm = self.ln2(C_u) 

149 P_k_t = torch.bmm(w3_x_u_norm, C_u_norm.transpose(1, 2)) 

150 P_k_t_b = F.softmax(P_k_t, dim=2) 

151 P_k_t_b_t = P_k_t_b.transpose(1, 2) 

152 

153 # attention weighting 

154 a_k = x_u.unsqueeze(1).repeat(1, self.k, 1, 1).matmul(self.w_k_1) 

155 P_t_k = F.softmax( 

156 torch.tanh(a_k).matmul(self.w_k_2.reshape(self.k, self.embedding_size, 1)).squeeze(3), 

157 dim=2, 

158 ) 

159 

160 # interest embedding generation 

161 mul_p = P_k_t_b_t.mul(P_t_k) 

162 x_u_re = x_u.unsqueeze(1).repeat(1, self.k, 1, 1) 

163 mul_p_re = mul_p.unsqueeze(3) 

164 delta_k = x_u_re.mul(mul_p_re).sum(2) 

165 delta_k = F.normalize(delta_k, p=2, dim=2) 

166 

167 # prototype sequence 

168 x_u_bar = P_k_t_b.matmul(C_u) 

169 C_apt = F.softmax(torch.tanh(x_u_bar.matmul(self.w3)).matmul(self.w4), dim=1) 

170 C_apt = C_apt.reshape(-1, 1, self.max_seq_length).matmul(x_u_bar) 

171 C_apt = self.ln4(C_apt) 

172 

173 # aggregation weight 

174 e_k = delta_k.bmm(C_apt.reshape(-1, self.embedding_size, 1)) / self.tau 

175 e_k_u = F.softmax(e_k.squeeze(2), dim=1) 

176 v_u = e_k_u.unsqueeze(2).mul(delta_k).sum(dim=1) 

177 

178 return v_u 

179 

180 def full_sort_predict(self, interaction): 

181 item_seq = interaction[self.ITEM_SEQ] 

182 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

183 seq_output = self.forward(item_seq, item_seq_len) 

184 test_items_emb = self.item_embedding.weight 

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

186 return scores