radames/Text2Human-API
1
1def accuracy(pred, target, topk=1, thresh=None):2 """Calculate accuracy according to the prediction and target.3 4 Args:5 pred (torch.Tensor): The model prediction, shape (N, num_class, ...)6 target (torch.Tensor): The target of each prediction, shape (N, , ...)7 topk (int | tuple[int], optional): If the predictions in ``topk``8 matches the target, the predictions will be regarded as9 correct ones. Defaults to 1.10 thresh (float, optional): If not None, predictions with scores under11 this threshold are considered incorrect. Default to None.12 13 Returns:14 float | tuple[float]: If the input ``topk`` is a single integer,15 the function will return a single float as accuracy. If16 ``topk`` is a tuple containing multiple integers, the17 function will return a tuple containing accuracies of18 each ``topk`` number.19 """20 assert isinstance(topk, (int, tuple))21 if isinstance(topk, int):22 topk = (topk, )23 return_single = True24 else:25 return_single = False26 27 maxk = max(topk)28 if pred.size(0) == 0:29 accu = [pred.new_tensor(0.) for i in range(len(topk))]30 return accu[0] if return_single else accu31 assert pred.ndim == target.ndim + 132 assert pred.size(0) == target.size(0)33 assert maxk <= pred.size(1), \34 f'maxk {maxk} exceeds pred dimension {pred.size(1)}'35 pred_value, pred_label = pred.topk(maxk, dim=1)36 # transpose to shape (maxk, N, ...)37 pred_label = pred_label.transpose(0, 1)38 correct = pred_label.eq(target.unsqueeze(0).expand_as(pred_label))39 if thresh is not None:40 # Only prediction values larger than thresh are counted as correct41 correct = correct & (pred_value > thresh).t()42 res = []43 for k in topk:44 correct_k = correct[:k].view(-1).float().sum(0, keepdim=True)45 res.append(correct_k.mul_(100.0 / target.numel()))46 return res[0] if return_single else res47 