(predictions, targets, k)
| 95 | |
| 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): |
| 107 | cnt += solve(torch.from_numpy(pred), targets[index], k) |
| 108 | |
| 109 | return cnt * 100.0 / len(predictions) |
| 110 | |
| 111 | |
| 112 | def segmentation_metrics(predictions, targets, classes): |
no test coverage detected