| 212 | return supervised_output |
| 213 | |
| 214 | class 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) |