Coverage for hopwise/model/general_recommender/fism.py: 96%

96 statements  

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

1# @Time : 2020/09/28 

2# @Author : Kaiyuan Li 

3# @email : tsotfsk@outlook.com 

4 

5"""FISM 

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

7Reference: 

8 S. Kabbur et al. "FISM: Factored item similarity models for top-n recommender systems" in KDD 2013 

9 

10Reference code: 

11 https://github.com/AaronHeee/Neural-Attentive-Item-Similarity-Model 

12""" 

13 

14import torch 

15from torch import nn 

16from torch.nn.init import normal_ 

17 

18from hopwise.model.abstract_recommender import GeneralRecommender 

19from hopwise.utils import InputType 

20 

21 

22class FISM(GeneralRecommender): 

23 """FISM is an item-based model for generating top-N recommendations that learns the 

24 item-item similarity matrix as the product of two low dimensional latent factor matrices. 

25 These matrices are learned using a structural equation modeling approach, where in the 

26 value being estimated is not used for its own estimation. 

27 

28 """ 

29 

30 input_type = InputType.POINTWISE 

31 

32 def __init__(self, config, dataset): 

33 super().__init__(config, dataset) 

34 

35 # load dataset info 

36 self.LABEL = config["LABEL_FIELD"] 

37 # get all users' history interaction information.the history item 

38 # matrix is padding by the maximum number of a user's interactions 

39 ( 

40 self.history_item_matrix, 

41 self.history_lens, 

42 self.mask_mat, 

43 ) = self.get_history_info(dataset) 

44 

45 # load parameters info 

46 self.embedding_size = config["embedding_size"] 

47 self.reg_weights = config["reg_weights"] 

48 self.alpha = config["alpha"] 

49 self.split_to = config["split_to"] 

50 

51 # split the too large dataset into the specified pieces 

52 if self.split_to > 0: 

53 self.group = torch.chunk(torch.arange(self.n_items).to(self.device), self.split_to) 

54 else: 

55 self.logger.warning( 

56 "Pay Attetion!! the `split_to` is set to 0. If you catch a OMM error in this case, " 

57 + "you need to increase it \n\t\t\tuntil the error disappears. For example, " 

58 + "you can append it in the command line such as `--split_to=5`" 

59 ) 

60 

61 # define layers and loss 

62 # construct source and destination item embedding matrix 

63 self.item_src_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0) 

64 self.item_dst_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0) 

65 self.user_bias = nn.Parameter(torch.zeros(self.n_users)) 

66 self.item_bias = nn.Parameter(torch.zeros(self.n_items)) 

67 self.bceloss = nn.BCEWithLogitsLoss() 

68 

69 # parameters initialization 

70 self.apply(self._init_weights) 

71 

72 def get_history_info(self, dataset): 

73 """Get the user history interaction information 

74 

75 Args: 

76 dataset (DataSet): train dataset 

77 

78 Returns: 

79 tuple: (history_item_matrix, history_lens, mask_mat) 

80 

81 """ 

82 history_item_matrix, _, history_lens = dataset.history_item_matrix() 

83 history_item_matrix = history_item_matrix.to(self.device) 

84 history_lens = history_lens.to(self.device) 

85 arange_tensor = torch.arange(history_item_matrix.shape[1]).to(self.device) 

86 mask_mat = (arange_tensor < history_lens.unsqueeze(1)).float() 

87 return history_item_matrix, history_lens, mask_mat 

88 

89 def reg_loss(self): 

90 """Calculate the reg loss for embedding layers 

91 

92 Returns: 

93 torch.Tensor: reg loss 

94 

95 """ 

96 reg_1, reg_2 = self.reg_weights 

97 loss_1 = reg_1 * self.item_src_embedding.weight.norm(2) 

98 loss_2 = reg_2 * self.item_dst_embedding.weight.norm(2) 

99 

100 return loss_1 + loss_2 

101 

102 def _init_weights(self, module): 

103 """Initialize the module's parameters 

104 

105 Note: 

106 It's a little different from the source code, because pytorch has no function to initialize 

107 the parameters by truncated normal distribution, so we replace it with xavier normal distribution 

108 

109 """ 

110 if isinstance(module, nn.Embedding): 

111 normal_(module.weight.data, 0, 0.01) 

112 

113 def inter_forward(self, user, item): 

114 """Forward the model by interaction""" 

115 user_inter = self.history_item_matrix[user] 

116 item_num = self.history_lens[user].unsqueeze(1) 

117 batch_mask_mat = self.mask_mat[user] 

118 user_history = self.item_src_embedding(user_inter) # batch_size x max_len x embedding_size 

119 target = self.item_dst_embedding(item) # batch_size x embedding_size 

120 user_bias = self.user_bias[user] # batch_size x 1 

121 item_bias = self.item_bias[item] 

122 similarity = torch.bmm(user_history, target.unsqueeze(2)).squeeze(2) # batch_size x max_len 

123 similarity = batch_mask_mat * similarity 

124 coeff = torch.pow(item_num.squeeze(1), -self.alpha) 

125 scores = torch.sigmoid(coeff.float() * torch.sum(similarity, dim=1) + user_bias + item_bias) 

126 return scores 

127 

128 def user_forward(self, user_input, item_num, user_bias, repeats=None, pred_slc=None): 

129 """Forward the model by user 

130 

131 Args: 

132 user_input (torch.Tensor): user input tensor 

133 item_num (torch.Tensor): user history interaction lens 

134 repeats (int, optional): the number of items to be evaluated 

135 pred_slc (torch.Tensor, optional): continuous index which controls the current evaluation items, 

136 if pred_slc is None, it will evaluate all items 

137 

138 Returns: 

139 torch.Tensor: result 

140 

141 """ 

142 item_num = item_num.repeat(repeats, 1) 

143 user_history = self.item_src_embedding(user_input) # inter_num x embedding_size 

144 user_history = user_history.repeat(repeats, 1, 1) # target_items x inter_num x embedding_size 

145 if pred_slc is None: 

146 targets = self.item_dst_embedding.weight # target_items x embedding_size 

147 item_bias = self.item_bias 

148 else: 

149 targets = self.item_dst_embedding(pred_slc) 

150 item_bias = self.item_bias[pred_slc] 

151 similarity = torch.bmm(user_history, targets.unsqueeze(2)).squeeze(2) # inter_num x target_items 

152 coeff = torch.pow(item_num.squeeze(1), -self.alpha) 

153 scores = coeff.float() * torch.sum(similarity, dim=1) + user_bias + item_bias 

154 return scores 

155 

156 def forward(self, user, item): 

157 return self.inter_forward(user, item) 

158 

159 def calculate_loss(self, interaction): 

160 user = interaction[self.USER_ID] 

161 item = interaction[self.ITEM_ID] 

162 label = interaction[self.LABEL] 

163 output = self.forward(user, item) 

164 loss = self.bceloss(output, label) + self.reg_loss() 

165 return loss 

166 

167 def full_sort_predict(self, interaction): 

168 user = interaction[self.USER_ID] 

169 batch_user_bias = self.user_bias[user] 

170 user_inters = self.history_item_matrix[user] 

171 item_nums = self.history_lens[user] 

172 scores = [] 

173 

174 # test users one by one, if the number of items is too large, we will split it to some pieces 

175 for user_input, item_num, user_bias in zip(user_inters, item_nums.unsqueeze(1), batch_user_bias): 

176 if self.split_to <= 0: 

177 output = self.user_forward(user_input[:item_num], item_num, user_bias, repeats=self.n_items) 

178 else: 

179 output = [] 

180 for mask in self.group: 

181 tmp_output = self.user_forward( 

182 user_input[:item_num], 

183 item_num, 

184 user_bias, 

185 repeats=len(mask), 

186 pred_slc=mask, 

187 ) 

188 output.append(tmp_output) 

189 output = torch.cat(output, dim=0) 

190 scores.append(output) 

191 result = torch.cat(scores, dim=0) 

192 return result 

193 

194 def predict(self, interaction): 

195 user = interaction[self.USER_ID] 

196 item = interaction[self.ITEM_ID] 

197 output = torch.sigmoid(self.forward(user, item)) 

198 return output