import torch def predict(model, q_ids, q_mask, opt_ids, opt_mask, device): model.eval() q_ids = q_ids.to(device) q_mask = q_mask.to(device) opt_ids = opt_ids.to(device) opt_mask = opt_mask.to(device) with torch.no_grad(): scores = model( q_ids, q_mask, opt_ids, opt_mask ) top3 = torch.topk( scores, k=3, dim=1 ).indices return scores, top3