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

1# @Time : 2020/09/01 

2# @Author : Kaiyuan Li 

3# @email : tsotfsk@outlook.com 

4 

5# UPDATE: 

6# @Time : 2020/10/14 

7# @Author : Kaiyuan Li 

8# @Email : tsotfsk@outlook.com 

9 

10"""NAIS 

11###################################### 

12Reference: 

13 Xiangnan He et al. "NAIS: Neural Attentive Item Similarity Model for Recommendation." in TKDE 2018. 

14 

15Reference code: 

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

17""" 

18 

19import torch 

20from torch import nn 

21from torch.nn.init import constant_, normal_, xavier_normal_ 

22 

23from hopwise.model.abstract_recommender import GeneralRecommender 

24from hopwise.model.layers import MLPLayers 

25from hopwise.utils import InputType 

26 

27 

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. 

32 

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. 

36 

37 """ 

38 

39 input_type = InputType.POINTWISE 

40 

41 def __init__(self, config, dataset): 

42 super().__init__(config, dataset) 

43 

44 # load dataset info 

45 self.LABEL = config["LABEL_FIELD"] 

46 

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) 

54 

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"] 

64 

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 ) 

75 

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() 

89 

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) 

97 

98 def _init_weights(self, module): 

99 """Initialize the module's parameters 

100 

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 

104 

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) 

112 

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) 

123 

124 def get_history_info(self, dataset): 

125 """Get the user history interaction information 

126 

127 Args: 

128 dataset (DataSet): train dataset 

129 

130 Returns: 

131 tuple: (history_item_matrix, history_lens, mask_mat) 

132 

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 

140 

141 def reg_loss(self): 

142 """Calculate the reg loss for embedding layers and mlp layers 

143 

144 Returns: 

145 torch.Tensor: reg loss 

146 

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 

156 

157 def attention_mlp(self, inter, target): 

158 """Layers of attention which support `prod` and `concat` 

159 

160 Args: 

161 inter (torch.Tensor): the embedding of history items 

162 target (torch.Tensor): the embedding of target items 

163 

164 Returns: 

165 torch.Tensor: the result of attention 

166 

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 

175 

176 logits = torch.matmul(mlp_output, self.weight_layer).squeeze(2) # batch_size x max_len 

177 return logits 

178 

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 

181 

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 

188 

189 Returns: 

190 torch.Tensor: final output 

191 

192 """ 

193 exp_logits = torch.exp(logits) # batch_size x max_len 

194 

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) 

199 

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

201 output = coeff.float() * torch.sum(weights * similarity, dim=1) + bias 

202 

203 return output 

204 

205 def softmax(self, similarity, logits, item_num, bias): 

206 """Softmax the user history features and get the final output 

207 

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 

213 

214 Returns: 

215 torch.Tensor: final output 

216 

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) 

224 

225 return output 

226 

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 

239 

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

241 """Forward the model by user 

242 

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 

249 

250 Returns: 

251 torch.Tensor: result 

252 

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 

267 

268 def forward(self, user, item): 

269 return self.inter_forward(user, item) 

270 

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 

278 

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 = [] 

284 

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 

303 

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