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
« 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
6r"""RepeatNet
7################################################
9Reference:
10 Pengjie Ren et al. "RepeatNet: A Repeat Aware Neural Recommendation Machine for Session-based Recommendation."
11 in AAAI 2019
13Reference code:
14 https://github.com/PengjieRen/RepeatNet.
16"""
18import torch
19from torch import nn
20from torch.nn import functional as F
21from torch.nn.init import constant_, xavier_normal_
23from hopwise.model.abstract_recommender import SequentialRecommender
24from hopwise.utils import InputType
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
32 """
34 input_type = InputType.POINTWISE
36 def __init__(self, config, dataset):
37 super().__init__(config, dataset)
39 # load the dataset information
40 self.device = config["device"]
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"]
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 )
72 self.loss_fct = F.nll_loss
74 # init the weight of the module
75 self.apply(self._init_weights)
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)
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
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
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
116 return prediction
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)
127 return loss
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()
144 return (-repeat_loss - explore_loss) / 2
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)
151 return prediction
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)
163 return scores
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)
179 def forward(self, all_memory, last_memory):
180 """Calculate the probability of Repeat and explore"""
181 all_memory_values = all_memory
183 all_memory = self.dropout(self.Ure(all_memory))
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)
189 output_ere = self.tanh(all_memory + last_memory)
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)
197 output_cre = self.Wcre(output_cre)
199 repeat_explore_mechanism = nn.Softmax(dim=-1)(output_cre)
201 return repeat_explore_mechanism
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)
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))
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)
225 output_er = self.tanh(last_memory + all_memory)
227 output_er = self.Vr(output_er).squeeze(2)
229 if mask is not None:
230 output_er.masked_fill_(mask, -1e9)
232 output_er = nn.Softmax(dim=-1)(output_er)
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)
238 return repeat_recommendation_decoder.to(self.device)
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)
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
259 all_memory = self.dropout(self.Ue(all_memory))
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)
265 output_ee = self.tanh(all_memory + last_memory)
266 output_ee = self.Ve(output_ee).squeeze(-1)
268 if mask is not None:
269 output_ee.masked_fill_(mask, -1e9)
271 output_ee = output_ee.unsqueeze(-1)
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))
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)
285 return explore_recommendation_decoder