(self, states, len_states)
| 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 |
nothing calls this directly
no outgoing calls
no test coverage detected