MCPcopy Create free account
hub / github.com/FunAudioLLM/FunMusic / th_accuracy

Function th_accuracy

inspiremusic/utils/common.py:71–90  ·  view source on GitHub ↗

Calculate accuracy. Args: pad_outputs (Tensor): Prediction tensors (B * Lmax, D). pad_targets (LongTensor): Target label tensors (B, Lmax). ignore_label (int): Ignore label id. Returns: torch.Tensor: Accuracy value (0.0 - 1.0).

(pad_outputs: torch.Tensor, pad_targets: torch.Tensor,
                ignore_label: int)

Source from the content-addressed store, hash-verified

69
70
71def th_accuracy(pad_outputs: torch.Tensor, pad_targets: torch.Tensor,
72 ignore_label: int) -> torch.Tensor:
73 """Calculate accuracy.
74
75 Args:
76 pad_outputs (Tensor): Prediction tensors (B * Lmax, D).
77 pad_targets (LongTensor): Target label tensors (B, Lmax).
78 ignore_label (int): Ignore label id.
79
80 Returns:
81 torch.Tensor: Accuracy value (0.0 - 1.0).
82
83 """
84 pad_pred = pad_outputs.view(pad_targets.size(0), pad_targets.size(1),
85 pad_outputs.size(1)).argmax(2)
86 mask = pad_targets != ignore_label
87 numerator = torch.sum(
88 pad_pred.masked_select(mask) == pad_targets.masked_select(mask))
89 denominator = torch.sum(mask)
90 return (numerator / denominator).detach()
91
92def get_padding(kernel_size, dilation=1):
93 return int((kernel_size * dilation - dilation) / 2)

Callers 1

forwardMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected