Coverage for hopwise/model/general_recommender/nais.py: 90%
149 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/01
2# @Author : Kaiyuan Li
3# @email : tsotfsk@outlook.com
5# UPDATE:
6# @Time : 2020/10/14
7# @Author : Kaiyuan Li
8# @Email : tsotfsk@outlook.com
10"""NAIS
11######################################
12Reference:
13 Xiangnan He et al. "NAIS: Neural Attentive Item Similarity Model for Recommendation." in TKDE 2018.
15Reference code:
16 https://github.com/AaronHeee/Neural-Attentive-Item-Similarity-Model
17"""
19import torch
20from torch import nn
21from torch.nn.init import constant_, normal_, xavier_normal_
23from hopwise.model.abstract_recommender import GeneralRecommender
24from hopwise.model.layers import MLPLayers
25from hopwise.utils import InputType
28class NAIS(GeneralRecommender):
29 """NAIS is an attention network, which is capable of distinguishing which historical items
30 in a user profile are more important for a prediction. We just implement the model following
31 the original author with a pointwise training mode.
33 Note:
34 instead of forming a minibatch as all training instances of a randomly sampled user which is
35 mentioned in the original paper, we still train the model by a randomly sampled interactions.
37 """
39 input_type = InputType.POINTWISE
41 def __init__(self, config, dataset):
42 super().__init__(config, dataset)
44 # load dataset info
45 self.LABEL = config["LABEL_FIELD"]
47 # get all users' history interaction information.the history item
48 # matrix is padding by the maximum number of a user's interactions
49 (
50 self.history_item_matrix,
51 self.history_lens,
52 self.mask_mat,
53 ) = self.get_history_info(dataset)
55 # load parameters info
56 self.embedding_size = config["embedding_size"]
57 self.weight_size = config["weight_size"]
58 self.algorithm = config["algorithm"]
59 self.reg_weights = config["reg_weights"]
60 self.alpha = config["alpha"]
61 self.beta = config["beta"]
62 self.split_to = config["split_to"]
63 self.pretrain_path = config["pretrain_path"]
65 # split the too large dataset into the specified pieces
66 if self.split_to > 0:
67 self.logger.info(f"split the n_items to {self.split_to} pieces")
68 self.group = torch.chunk(torch.arange(self.n_items).to(self.device), self.split_to)
69 else:
70 self.logger.warning(
71 "Pay Attetion!! the `split_to` is set to 0. If you catch a OMM error in this case, "
72 + "you need to increase it \n\t\t\tuntil the error disappears. For example, "
73 + "you can append it in the command line such as `--split_to=5`"
74 )
76 # define layers and loss
77 # construct source and destination item embedding matrix
78 self.item_src_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0)
79 self.item_dst_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0)
80 self.bias = nn.Parameter(torch.zeros(self.n_items))
81 if self.algorithm == "concat":
82 self.mlp_layers = MLPLayers([self.embedding_size * 2, self.weight_size])
83 elif self.algorithm == "prod":
84 self.mlp_layers = MLPLayers([self.embedding_size, self.weight_size])
85 else:
86 raise ValueError(f"NAIS just support attention type in ['concat', 'prod'] but get {self.algorithm}")
87 self.weight_layer = nn.Parameter(torch.ones(self.weight_size, 1))
88 self.bceloss = nn.BCEWithLogitsLoss()
90 # parameters initialization
91 if self.pretrain_path is not None:
92 self.logger.info(f"use pretrain from [{self.pretrain_path}]...")
93 self._load_pretrain()
94 else:
95 self.logger.info("unused pretrain...")
96 self.apply(self._init_weights)
98 def _init_weights(self, module):
99 """Initialize the module's parameters
101 Note:
102 It's a little different from the source code, because pytorch has no function to initialize
103 the parameters by truncated normal distribution, so we replace it with xavier normal distribution
105 """
106 if isinstance(module, nn.Embedding):
107 normal_(module.weight.data, 0, 0.01)
108 elif isinstance(module, nn.Linear):
109 xavier_normal_(module.weight.data)
110 if module.bias is not None:
111 constant_(module.bias.data, 0)
113 def _load_pretrain(self):
114 """A simple implementation of loading pretrained parameters."""
115 fism = torch.load(self.pretrain_path)["state_dict"]
116 self.item_src_embedding.weight.data.copy_(fism["item_src_embedding.weight"])
117 self.item_dst_embedding.weight.data.copy_(fism["item_dst_embedding.weight"])
118 for name, parm in self.mlp_layers.named_parameters():
119 if name.endswith("weight"):
120 xavier_normal_(parm.data)
121 elif name.endswith("bias"):
122 constant_(parm.data, 0)
124 def get_history_info(self, dataset):
125 """Get the user history interaction information
127 Args:
128 dataset (DataSet): train dataset
130 Returns:
131 tuple: (history_item_matrix, history_lens, mask_mat)
133 """
134 history_item_matrix, _, history_lens = dataset.history_item_matrix()
135 history_item_matrix = history_item_matrix.to(self.device)
136 history_lens = history_lens.to(self.device)
137 arange_tensor = torch.arange(history_item_matrix.shape[1]).to(self.device)
138 mask_mat = (arange_tensor < history_lens.unsqueeze(1)).float()
139 return history_item_matrix, history_lens, mask_mat
141 def reg_loss(self):
142 """Calculate the reg loss for embedding layers and mlp layers
144 Returns:
145 torch.Tensor: reg loss
147 """
148 reg_1, reg_2, reg_3 = self.reg_weights
149 loss_1 = reg_1 * self.item_src_embedding.weight.norm(2)
150 loss_2 = reg_2 * self.item_dst_embedding.weight.norm(2)
151 loss_3 = 0
152 for name, parm in self.mlp_layers.named_parameters():
153 if name.endswith("weight"):
154 loss_3 = loss_3 + reg_3 * parm.norm(2)
155 return loss_1 + loss_2 + loss_3
157 def attention_mlp(self, inter, target):
158 """Layers of attention which support `prod` and `concat`
160 Args:
161 inter (torch.Tensor): the embedding of history items
162 target (torch.Tensor): the embedding of target items
164 Returns:
165 torch.Tensor: the result of attention
167 """
168 if self.algorithm == "prod":
169 mlp_input = inter * target.unsqueeze(1) # batch_size x max_len x embedding_size
170 else:
171 mlp_input = torch.cat(
172 [inter, target.unsqueeze(1).expand_as(inter)], dim=2
173 ) # batch_size x max_len x embedding_size*2
174 mlp_output = self.mlp_layers(mlp_input) # batch_size x max_len x weight_size
176 logits = torch.matmul(mlp_output, self.weight_layer).squeeze(2) # batch_size x max_len
177 return logits
179 def mask_softmax(self, similarity, logits, bias, item_num, batch_mask_mat):
180 """Softmax the unmasked user history items and get the final output
182 Args:
183 similarity (torch.Tensor): the similarity between the history items and target items
184 logits (torch.Tensor): the initial weights of the history items
185 item_num (torch.Tensor): user history interaction lengths
186 bias (torch.Tensor): bias
187 batch_mask_mat (torch.Tensor): the mask of user history interactions
189 Returns:
190 torch.Tensor: final output
192 """
193 exp_logits = torch.exp(logits) # batch_size x max_len
195 exp_logits = batch_mask_mat * exp_logits # batch_size x max_len
196 exp_sum = torch.sum(exp_logits, dim=1, keepdim=True)
197 exp_sum = torch.pow(exp_sum, self.beta)
198 weights = torch.div(exp_logits, exp_sum)
200 coeff = torch.pow(item_num.squeeze(1), -self.alpha)
201 output = coeff.float() * torch.sum(weights * similarity, dim=1) + bias
203 return output
205 def softmax(self, similarity, logits, item_num, bias):
206 """Softmax the user history features and get the final output
208 Args:
209 similarity (torch.Tensor): the similarity between the history items and target items
210 logits (torch.Tensor): the initial weights of the history items
211 item_num (torch.Tensor): user history interaction lengths
212 bias (torch.Tensor): bias
214 Returns:
215 torch.Tensor: final output
217 """
218 exp_logits = torch.exp(logits) # batch_size x max_len
219 exp_sum = torch.sum(exp_logits, dim=1, keepdim=True)
220 exp_sum = torch.pow(exp_sum, self.beta)
221 weights = torch.div(exp_logits, exp_sum)
222 coeff = torch.pow(item_num.squeeze(1), -self.alpha)
223 output = torch.sigmoid(coeff.float() * torch.sum(weights * similarity, dim=1) + bias)
225 return output
227 def inter_forward(self, user, item):
228 """Forward the model by interaction"""
229 user_inter = self.history_item_matrix[user]
230 item_num = self.history_lens[user].unsqueeze(1)
231 batch_mask_mat = self.mask_mat[user]
232 user_history = self.item_src_embedding(user_inter) # batch_size x max_len x embedding_size
233 target = self.item_dst_embedding(item) # batch_size x embedding_size
234 bias = self.bias[item] # batch_size x 1
235 similarity = torch.bmm(user_history, target.unsqueeze(2)).squeeze(2) # batch_size x max_len
236 logits = self.attention_mlp(user_history, target)
237 scores = self.mask_softmax(similarity, logits, bias, item_num, batch_mask_mat)
238 return scores
240 def user_forward(self, user_input, item_num, repeats=None, pred_slc=None):
241 """Forward the model by user
243 Args:
244 user_input (torch.Tensor): user input tensor
245 item_num (torch.Tensor): user history interaction lens
246 repeats (int, optional): the number of items to be evaluated
247 pred_slc (torch.Tensor, optional): continuous index which controls the current evaluation items,
248 if pred_slc is None, it will evaluate all items
250 Returns:
251 torch.Tensor: result
253 """
254 item_num = item_num.repeat(repeats, 1)
255 user_history = self.item_src_embedding(user_input) # inter_num x embedding_size
256 user_history = user_history.repeat(repeats, 1, 1) # target_items x inter_num x embedding_size
257 if pred_slc is None:
258 targets = self.item_dst_embedding.weight # target_items x embedding_size
259 bias = self.bias
260 else:
261 targets = self.item_dst_embedding(pred_slc)
262 bias = self.bias[pred_slc]
263 similarity = torch.bmm(user_history, targets.unsqueeze(2)).squeeze(2) # inter_num x target_items
264 logits = self.attention_mlp(user_history, targets)
265 scores = self.softmax(similarity, logits, item_num, bias)
266 return scores
268 def forward(self, user, item):
269 return self.inter_forward(user, item)
271 def calculate_loss(self, interaction):
272 user = interaction[self.USER_ID]
273 item = interaction[self.ITEM_ID]
274 label = interaction[self.LABEL]
275 output = self.forward(user, item)
276 loss = self.bceloss(output, label) + self.reg_loss()
277 return loss
279 def full_sort_predict(self, interaction):
280 user = interaction[self.USER_ID]
281 user_inters = self.history_item_matrix[user]
282 item_nums = self.history_lens[user]
283 scores = []
285 # test users one by one, if the number of items is too large, we will split it to some pieces
286 for user_input, item_num in zip(user_inters, item_nums.unsqueeze(1)):
287 if self.split_to <= 0:
288 output = self.user_forward(user_input[:item_num], item_num, repeats=self.n_items)
289 else:
290 output = []
291 for mask in self.group:
292 tmp_output = self.user_forward(
293 user_input[:item_num],
294 item_num,
295 repeats=len(mask),
296 pred_slc=mask,
297 )
298 output.append(tmp_output)
299 output = torch.cat(output, dim=0)
300 scores.append(output)
301 result = torch.cat(scores, dim=0)
302 return result
304 def predict(self, interaction):
305 user = interaction[self.USER_ID]
306 item = interaction[self.ITEM_ID]
307 output = torch.sigmoid(self.forward(user, item))
308 return output