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
« 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
5"""FISM
6#######################################
7Reference:
8 S. Kabbur et al. "FISM: Factored item similarity models for top-n recommender systems" in KDD 2013
10Reference code:
11 https://github.com/AaronHeee/Neural-Attentive-Item-Similarity-Model
12"""
14import torch
15from torch import nn
16from torch.nn.init import normal_
18from hopwise.model.abstract_recommender import GeneralRecommender
19from hopwise.utils import InputType
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.
28 """
30 input_type = InputType.POINTWISE
32 def __init__(self, config, dataset):
33 super().__init__(config, dataset)
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)
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"]
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 )
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()
69 # parameters initialization
70 self.apply(self._init_weights)
72 def get_history_info(self, dataset):
73 """Get the user history interaction information
75 Args:
76 dataset (DataSet): train dataset
78 Returns:
79 tuple: (history_item_matrix, history_lens, mask_mat)
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
89 def reg_loss(self):
90 """Calculate the reg loss for embedding layers
92 Returns:
93 torch.Tensor: reg loss
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)
100 return loss_1 + loss_2
102 def _init_weights(self, module):
103 """Initialize the module's parameters
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
109 """
110 if isinstance(module, nn.Embedding):
111 normal_(module.weight.data, 0, 0.01)
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
128 def user_forward(self, user_input, item_num, user_bias, repeats=None, pred_slc=None):
129 """Forward the model by user
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
138 Returns:
139 torch.Tensor: result
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
156 def forward(self, user, item):
157 return self.inter_forward(user, item)
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
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 = []
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
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