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

Method forward_eval

sasrec.py:259–275  ·  view source on GitHub ↗
(self, states, len_states)

Source from the content-addressed store, hash-verified

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
278def evaluate_games(model, test_data, device, topk, save_logits=False, eval_type="test"):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected