convert text-index into text-label.
(self, text_index, text_prob=None, is_remove_duplicate=False)
| 18 | self.nclass = len(self.character) + 1 |
| 19 | |
| 20 | def decode(self, text_index, text_prob=None, is_remove_duplicate=False): |
| 21 | """convert text-index into text-label.""" |
| 22 | result_list = [] |
| 23 | ignored_tokens = self.get_ignored_tokens() |
| 24 | batch_size = len(text_index) |
| 25 | for batch_idx in range(batch_size): |
| 26 | selection = np.ones(len(text_index[batch_idx]), dtype=bool) |
| 27 | if is_remove_duplicate: |
| 28 | selection[1:] = text_index[batch_idx][1:] != text_index[ |
| 29 | batch_idx][:-1] |
| 30 | for ignored_token in ignored_tokens: |
| 31 | selection &= text_index[batch_idx] != ignored_token |
| 32 | |
| 33 | char_list = [ |
| 34 | self.character[text_id - 1] |
| 35 | for text_id in text_index[batch_idx][selection] |
| 36 | ] |
| 37 | if text_prob is not None: |
| 38 | conf_list = text_prob[batch_idx][selection] |
| 39 | else: |
| 40 | conf_list = [1] * len(selection) |
| 41 | if len(conf_list) == 0: |
| 42 | conf_list = [0] |
| 43 | |
| 44 | text = ''.join(char_list) |
| 45 | result_list.append((text, np.mean(conf_list).tolist())) |
| 46 | return result_list |
| 47 | |
| 48 | def __call__(self, preds, batch=None, *args, **kwargs): |
| 49 | if len(preds) == 2: # eval mode |
no test coverage detected