MCPcopy Create free account
hub / github.com/AkaliKong/MiniOneRec / SASRec

Class SASRec

sasrec.py:214–275  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

212 return supervised_output
213
214class SASRec(nn.Module):
215 def __init__(self, hidden_size, item_num, state_size, dropout, device, num_heads=1):
216 super(SASRec, self).__init__()
217 self.state_size = state_size
218 self.hidden_size = hidden_size
219 self.item_num = int(item_num)
220 self.dropout = nn.Dropout(dropout)
221 self.device = device
222 self.item_embeddings = nn.Embedding(
223 num_embeddings=item_num + 1,
224 embedding_dim=hidden_size,
225 )
226 nn.init.normal_(self.item_embeddings.weight, 0, 0.01)
227 self.positional_embeddings = nn.Embedding(
228 num_embeddings=state_size,
229 embedding_dim=hidden_size
230 )
231 # emb_dropout is added
232 self.emb_dropout = nn.Dropout(dropout)
233 self.ln_1 = nn.LayerNorm(hidden_size)
234 self.ln_2 = nn.LayerNorm(hidden_size)
235 self.ln_3 = nn.LayerNorm(hidden_size)
236 self.mh_attn = MultiHeadAttention(hidden_size, hidden_size, num_heads, dropout)
237 self.feed_forward = PositionwiseFeedForward(hidden_size, hidden_size, dropout)
238 self.s_fc = nn.Linear(hidden_size, item_num)
239 # self.ac_func = nn.ReLU()
240
241 def forward(self, states, len_states):
242 # inputs_emb = self.item_embeddings(states) * self.item_embeddings.embedding_dim ** 0.5
243 inputs_emb = self.item_embeddings(states)
244 inputs_emb += self.positional_embeddings(torch.arange(self.state_size).to(self.device))
245 seq = self.emb_dropout(inputs_emb)
246 mask = torch.ne(states, self.item_num).float().unsqueeze(-1).to(self.device)
247 seq *= mask
248 seq_normalized = self.ln_1(seq)
249 mh_attn_out = self.mh_attn(seq_normalized, seq)
250 ff_out = self.feed_forward(self.ln_2(mh_attn_out))
251 ff_out *= mask
252 ff_out = self.ln_3(ff_out)
253 # state_hidden = extract_axis_1(ff_out, len_states - 1)
254 indices = (len_states -1 ).view(-1, 1, 1).repeat(1, 1, self.hidden_size)
255 state_hidden = torch.gather(ff_out, 1, indices)
256 supervised_output = self.s_fc(state_hidden).squeeze()
257 return supervised_output
258
259 def forward_eval(self, states, len_states):
260 # inputs_emb = self.item_embeddings(states) * self.item_embeddings.embedding_dim ** 0.5
261 inputs_emb = self.item_embeddings(states)
262 inputs_emb += self.positional_embeddings(torch.arange(self.state_size).to(self.device))
263 seq = self.emb_dropout(inputs_emb)
264 mask = torch.ne(states, self.item_num).float().unsqueeze(-1).to(self.device)
265 seq *= mask
266 seq_normalized = self.ln_1(seq)
267 mh_attn_out = self.mh_attn(seq_normalized, seq)
268 ff_out = self.feed_forward(self.ln_2(mh_attn_out))
269 ff_out *= mask
270 ff_out = self.ln_3(ff_out)
271 # state_hidden = extract_axis_1(ff_out, len_states - 1)

Callers 3

trainFunction · 0.90
trainFunction · 0.90
mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected