(prob, target, k)
| 96 | |
| 97 | def topk_accuracy(predictions, targets, k): |
| 98 | def solve(prob, target, k): |
| 99 | _, indices = torch.topk(prob, k=k, sorted=True) |
| 100 | golden = torch.reshape(target, [-1, 1]) |
| 101 | correct = (golden == indices) * 1.0 |
| 102 | top_k_accuracy = torch.mean(correct) * k |
| 103 | return top_k_accuracy |
| 104 | |
| 105 | cnt = 0 |
| 106 | for index, pred in enumerate(predictions): |