Calculate top-k accuracy.
(output: torch.Tensor, target: torch.Tensor, topk: Tuple[int, ...] = (1,))
| 310 | |
| 311 | |
| 312 | def accuracy(output: torch.Tensor, target: torch.Tensor, topk: Tuple[int, ...] = (1,)) -> List[float]: |
| 313 | """Calculate top-k accuracy.""" |
| 314 | pred = output.topk(max(topk), 1, True, True)[1].t() |
| 315 | correct = pred.eq(target.view(1, -1).expand_as(pred)) |
| 316 | return [float(correct[:k].reshape(-1).float().sum(0, keepdim=True).cpu().numpy()) for k in topk] |
| 317 | |
| 318 | |
| 319 | def get_input_dtype(precision: str) -> torch.dtype: |