Coverage for hopwise/model/sequential_recommender/repeatnet.py: 88%

161 statements  

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

1# @Time : 2020/11/22 8:30 

2# @Author : Shao Weiqi 

3# @Reviewer : Lin Kun, Fan xinyan 

4# @Email : shaoweiqi@ruc.edu.cn, xinyan.fan@ruc.edu.cn 

5 

6r"""RepeatNet 

7################################################ 

8 

9Reference: 

10 Pengjie Ren et al. "RepeatNet: A Repeat Aware Neural Recommendation Machine for Session-based Recommendation." 

11 in AAAI 2019 

12 

13Reference code: 

14 https://github.com/PengjieRen/RepeatNet. 

15 

16""" 

17 

18import torch 

19from torch import nn 

20from torch.nn import functional as F 

21from torch.nn.init import constant_, xavier_normal_ 

22 

23from hopwise.model.abstract_recommender import SequentialRecommender 

24from hopwise.utils import InputType 

25 

26 

27class RepeatNet(SequentialRecommender): 

28 r"""RepeatNet explores a hybrid encoder with an repeat module and explore module 

29 repeat module is used for finding out the repeat consume in sequential recommendation 

30 explore module is used for exploring new items for recommendation 

31 

32 """ 

33 

34 input_type = InputType.POINTWISE 

35 

36 def __init__(self, config, dataset): 

37 super().__init__(config, dataset) 

38 

39 # load the dataset information 

40 self.device = config["device"] 

41 

42 # load parameters 

43 self.embedding_size = config["embedding_size"] 

44 self.hidden_size = config["hidden_size"] 

45 self.joint_train = config["joint_train"] 

46 self.dropout_prob = config["dropout_prob"] 

47 

48 # define the layers and loss function 

49 self.item_matrix = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0) 

50 self.gru = nn.GRU(self.embedding_size, self.hidden_size, batch_first=True) 

51 self.repeat_explore_mechanism = Repeat_Explore_Mechanism( 

52 self.device, 

53 hidden_size=self.hidden_size, 

54 seq_len=self.max_seq_length, 

55 dropout_prob=self.dropout_prob, 

56 ) 

57 self.repeat_recommendation_decoder = Repeat_Recommendation_Decoder( 

58 self.device, 

59 hidden_size=self.hidden_size, 

60 seq_len=self.max_seq_length, 

61 num_item=self.n_items, 

62 dropout_prob=self.dropout_prob, 

63 ) 

64 self.explore_recommendation_decoder = Explore_Recommendation_Decoder( 

65 hidden_size=self.hidden_size, 

66 seq_len=self.max_seq_length, 

67 num_item=self.n_items, 

68 device=self.device, 

69 dropout_prob=self.dropout_prob, 

70 ) 

71 

72 self.loss_fct = F.nll_loss 

73 

74 # init the weight of the module 

75 self.apply(self._init_weights) 

76 

77 def _init_weights(self, module): 

78 if isinstance(module, nn.Embedding): 

79 xavier_normal_(module.weight.data) 

80 elif isinstance(module, nn.Linear): 

81 xavier_normal_(module.weight.data) 

82 if module.bias is not None: 

83 constant_(module.bias.data, 0) 

84 

85 def forward(self, item_seq, item_seq_len): 

86 batch_seq_item_embedding = self.item_matrix(item_seq) 

87 # batch_size * seq_len == embedding ==>> batch_size * seq_len * embedding_size 

88 

89 all_memory, _ = self.gru(batch_seq_item_embedding) 

90 last_memory = self.gather_indexes(all_memory, item_seq_len - 1) 

91 # all_memory: batch_size * item_seq * hidden_size 

92 # last_memory: batch_size * hidden_size 

93 timeline_mask = item_seq == 0 

94 

95 self.repeat_explore = self.repeat_explore_mechanism.forward(all_memory=all_memory, last_memory=last_memory) 

96 # batch_size * 2 

97 repeat_recommendation_decoder = self.repeat_recommendation_decoder.forward( 

98 all_memory=all_memory, 

99 last_memory=last_memory, 

100 item_seq=item_seq, 

101 mask=timeline_mask, 

102 ) 

103 # batch_size * num_item 

104 explore_recommendation_decoder = self.explore_recommendation_decoder.forward( 

105 all_memory=all_memory, 

106 last_memory=last_memory, 

107 item_seq=item_seq, 

108 mask=timeline_mask, 

109 ) 

110 # batch_size * num_item 

111 prediction = repeat_recommendation_decoder * self.repeat_explore[:, 0].unsqueeze( 

112 1 

113 ) + explore_recommendation_decoder * self.repeat_explore[:, 1].unsqueeze(1) 

114 # batch_size * num_item 

115 

116 return prediction 

117 

118 def calculate_loss(self, interaction): 

119 item_seq = interaction[self.ITEM_SEQ] 

120 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

121 pos_item = interaction[self.POS_ITEM_ID] 

122 prediction = self.forward(item_seq, item_seq_len) 

123 loss = self.loss_fct((prediction + 1e-8).log(), pos_item, ignore_index=0) 

124 if self.joint_train is True: 

125 loss += self.repeat_explore_loss(item_seq, pos_item) 

126 

127 return loss 

128 

129 def repeat_explore_loss(self, item_seq, pos_item): 

130 batch_size = item_seq.size(0) 

131 repeat, explore = ( 

132 torch.zeros(batch_size).to(self.device), 

133 torch.ones(batch_size).to(self.device), 

134 ) 

135 index = 0 

136 for seq_item_ex, pos_item_ex in zip(item_seq, pos_item): 

137 if pos_item_ex in seq_item_ex: 

138 repeat[index] = 1 

139 explore[index] = 0 

140 index += 1 

141 repeat_loss = torch.mul(repeat.unsqueeze(1), torch.log(self.repeat_explore[:, 0] + 1e-8)).mean() 

142 explore_loss = torch.mul(explore.unsqueeze(1), torch.log(self.repeat_explore[:, 1] + 1e-8)).mean() 

143 

144 return (-repeat_loss - explore_loss) / 2 

145 

146 def full_sort_predict(self, interaction): 

147 item_seq = interaction[self.ITEM_SEQ] 

148 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

149 prediction = self.forward(item_seq, item_seq_len) 

150 

151 return prediction 

152 

153 def predict(self, interaction): 

154 item_seq = interaction[self.ITEM_SEQ] 

155 test_item = interaction[self.ITEM_ID] 

156 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

157 seq_output = self.forward(item_seq, item_seq_len) 

158 # batch_size * num_items 

159 seq_output = seq_output.unsqueeze(-1) 

160 # batch_size * num_items * 1 

161 scores = self.gather_indexes(seq_output, test_item).squeeze(-1) 

162 

163 return scores 

164 

165 

166class Repeat_Explore_Mechanism(nn.Module): 

167 def __init__(self, device, hidden_size, seq_len, dropout_prob): 

168 super().__init__() 

169 self.dropout = nn.Dropout(dropout_prob) 

170 self.hidden_size = hidden_size 

171 self.device = device 

172 self.seq_len = seq_len 

173 self.Wre = nn.Linear(hidden_size, hidden_size, bias=False) 

174 self.Ure = nn.Linear(hidden_size, hidden_size, bias=False) 

175 self.tanh = nn.Tanh() 

176 self.Vre = nn.Linear(hidden_size, 1, bias=False) 

177 self.Wcre = nn.Linear(hidden_size, 2, bias=False) 

178 

179 def forward(self, all_memory, last_memory): 

180 """Calculate the probability of Repeat and explore""" 

181 all_memory_values = all_memory 

182 

183 all_memory = self.dropout(self.Ure(all_memory)) 

184 

185 last_memory = self.dropout(self.Wre(last_memory)) 

186 last_memory = last_memory.unsqueeze(1) 

187 last_memory = last_memory.repeat(1, self.seq_len, 1) 

188 

189 output_ere = self.tanh(all_memory + last_memory) 

190 

191 output_ere = self.Vre(output_ere) 

192 alpha_are = nn.Softmax(dim=1)(output_ere) 

193 alpha_are = alpha_are.repeat(1, 1, self.hidden_size) 

194 output_cre = alpha_are * all_memory_values 

195 output_cre = output_cre.sum(dim=1) 

196 

197 output_cre = self.Wcre(output_cre) 

198 

199 repeat_explore_mechanism = nn.Softmax(dim=-1)(output_cre) 

200 

201 return repeat_explore_mechanism 

202 

203 

204class Repeat_Recommendation_Decoder(nn.Module): 

205 def __init__(self, device, hidden_size, seq_len, num_item, dropout_prob): 

206 super().__init__() 

207 self.dropout = nn.Dropout(dropout_prob) 

208 self.hidden_size = hidden_size 

209 self.device = device 

210 self.seq_len = seq_len 

211 self.num_item = num_item 

212 self.Wr = nn.Linear(hidden_size, hidden_size, bias=False) 

213 self.Ur = nn.Linear(hidden_size, hidden_size, bias=False) 

214 self.tanh = nn.Tanh() 

215 self.Vr = nn.Linear(hidden_size, 1) 

216 

217 def forward(self, all_memory, last_memory, item_seq, mask=None): 

218 """Calculate the the force of repeat""" 

219 all_memory = self.dropout(self.Ur(all_memory)) 

220 

221 last_memory = self.dropout(self.Wr(last_memory)) 

222 last_memory = last_memory.unsqueeze(1) 

223 last_memory = last_memory.repeat(1, self.seq_len, 1) 

224 

225 output_er = self.tanh(last_memory + all_memory) 

226 

227 output_er = self.Vr(output_er).squeeze(2) 

228 

229 if mask is not None: 

230 output_er.masked_fill_(mask, -1e9) 

231 

232 output_er = nn.Softmax(dim=-1)(output_er) 

233 

234 batch_size, b_len = item_seq.size() 

235 repeat_recommendation_decoder = torch.zeros([batch_size, self.num_item], device=self.device) 

236 repeat_recommendation_decoder.scatter_add_(1, item_seq, output_er) 

237 

238 return repeat_recommendation_decoder.to(self.device) 

239 

240 

241class Explore_Recommendation_Decoder(nn.Module): 

242 def __init__(self, hidden_size, seq_len, num_item, device, dropout_prob): 

243 super().__init__() 

244 self.dropout = nn.Dropout(dropout_prob) 

245 self.hidden_size = hidden_size 

246 self.seq_len = seq_len 

247 self.num_item = num_item 

248 self.device = device 

249 self.We = nn.Linear(hidden_size, hidden_size) 

250 self.Ue = nn.Linear(hidden_size, hidden_size) 

251 self.tanh = nn.Tanh() 

252 self.Ve = nn.Linear(hidden_size, 1) 

253 self.matrix_for_explore = nn.Linear(2 * self.hidden_size, self.num_item, bias=False) 

254 

255 def forward(self, all_memory, last_memory, item_seq, mask=None): 

256 """Calculate the force of explore""" 

257 all_memory_values, last_memory_values = all_memory, last_memory 

258 

259 all_memory = self.dropout(self.Ue(all_memory)) 

260 

261 last_memory = self.dropout(self.We(last_memory)) 

262 last_memory = last_memory.unsqueeze(1) 

263 last_memory = last_memory.repeat(1, self.seq_len, 1) 

264 

265 output_ee = self.tanh(all_memory + last_memory) 

266 output_ee = self.Ve(output_ee).squeeze(-1) 

267 

268 if mask is not None: 

269 output_ee.masked_fill_(mask, -1e9) 

270 

271 output_ee = output_ee.unsqueeze(-1) 

272 

273 alpha_e = nn.Softmax(dim=1)(output_ee) 

274 alpha_e = alpha_e.repeat(1, 1, self.hidden_size) 

275 output_e = (alpha_e * all_memory_values).sum(dim=1) 

276 output_e = torch.cat([output_e, last_memory_values], dim=1) 

277 output_e = self.dropout(self.matrix_for_explore(output_e)) 

278 

279 item_seq_first = item_seq[:, 0].unsqueeze(1).expand_as(item_seq) 

280 item_seq_first = item_seq_first.masked_fill(item_seq > 0, 0) 

281 item_seq_first.requires_grad_(False) 

282 output_e.scatter_add_(1, item_seq + item_seq_first, float("-inf") * torch.ones_like(item_seq)) 

283 explore_recommendation_decoder = nn.Softmax(1)(output_e) 

284 

285 return explore_recommendation_decoder