(self, states, len_states)
| 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) |
| 272 | indices = (len_states -1 ).view(-1, 1, 1).repeat(1, 1, self.hidden_size) |
| 273 | state_hidden = torch.gather(ff_out, 1, indices) |
| 274 | supervised_output = self.s_fc(state_hidden).squeeze() |
| 275 | return supervised_output |
| 276 | |
| 277 | |
| 278 | def evaluate_games(model, test_data, device, topk, save_logits=False, eval_type="test"): |
nothing calls this directly
no outgoing calls
no test coverage detected