Calculate sentence-level WER/CER score. :param torch.Tensor ys_hat: prediction (batch, seqlen) :param torch.Tensor ys_pad: reference (batch, seqlen) :param bool is_ctc: calculate CER score for CTC :return: sentence-level WER score :rtype float :return
(self, ys_hat, ys_pad, is_ctc=False)
| 123 | self.idx_space = None |
| 124 | |
| 125 | def __call__(self, ys_hat, ys_pad, is_ctc=False): |
| 126 | """Calculate sentence-level WER/CER score. |
| 127 | |
| 128 | :param torch.Tensor ys_hat: prediction (batch, seqlen) |
| 129 | :param torch.Tensor ys_pad: reference (batch, seqlen) |
| 130 | :param bool is_ctc: calculate CER score for CTC |
| 131 | :return: sentence-level WER score |
| 132 | :rtype float |
| 133 | :return: sentence-level CER score |
| 134 | :rtype float |
| 135 | """ |
| 136 | cer, wer = None, None |
| 137 | if is_ctc: |
| 138 | return self.calculate_cer_ctc(ys_hat, ys_pad) |
| 139 | elif not self.report_cer and not self.report_wer: |
| 140 | return cer, wer |
| 141 | |
| 142 | seqs_hat, seqs_true = self.convert_to_char(ys_hat, ys_pad) |
| 143 | if self.report_cer: |
| 144 | cer = self.calculate_cer(seqs_hat, seqs_true) |
| 145 | |
| 146 | if self.report_wer: |
| 147 | wer = self.calculate_wer(seqs_hat, seqs_true) |
| 148 | return cer, wer |
| 149 | |
| 150 | def calculate_cer_ctc(self, ys_hat, ys_pad): |
| 151 | """Calculate sentence-level CER score for CTC. |
nothing calls this directly
no test coverage detected