(self, hidden_size, item_num, state_size, dropout, device, num_heads=1)
| 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 |
nothing calls this directly
no test coverage detected