| 14 | use_space_char) |
| 15 | |
| 16 | def __call__(self, preds, batch=None, *args, **kwargs): |
| 17 | |
| 18 | if isinstance(preds, tuple): |
| 19 | if isinstance(preds[-1], dict): |
| 20 | preds = preds[-1]['align'][-1].detach().cpu().numpy() |
| 21 | else: |
| 22 | preds = preds[-1].detach().cpu().numpy() |
| 23 | if isinstance(preds, list): |
| 24 | preds = preds[-1].detach().cpu().numpy() |
| 25 | if isinstance(preds, torch.Tensor): |
| 26 | preds = preds.detach().cpu().numpy() |
| 27 | elif isinstance(preds, dict): |
| 28 | preds = preds['align'][-1].detach().cpu().numpy() |
| 29 | else: |
| 30 | preds = preds |
| 31 | preds_idx = preds.argmax(axis=2) |
| 32 | preds_prob = preds.max(axis=2) |
| 33 | text = self.decode(preds_idx, preds_prob, is_remove_duplicate=False) |
| 34 | if batch is None: |
| 35 | return text |
| 36 | label = batch[1] |
| 37 | label = self.decode(label) |
| 38 | return text, label |
| 39 | |
| 40 | def add_special_char(self, dict_character): |
| 41 | dict_character = ['</s>'] + dict_character |