Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
accuracy.py47 linesDownload Raw Back to losses
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