(self,
input_texts: Union[str, List[str]],
max_seq_len: int,
left_padding=False)
| 453 | return outputs, tokens |
| 454 | |
| 455 | def batch_encode(self, |
| 456 | input_texts: Union[str, List[str]], |
| 457 | max_seq_len: int, |
| 458 | left_padding=False): |
| 459 | if isinstance(input_texts, str): |
| 460 | input_texts = [input_texts] |
| 461 | tokens = [self.tokenizer(text) for text in input_texts] |
| 462 | max_len = min(max_seq_len, max([len(t) for t in tokens])) |
| 463 | for i in range(len(tokens)): |
| 464 | cur_input = tokens[i] |
| 465 | padding_len = max_len - len(cur_input) |
| 466 | if self.mode == 'none': |
| 467 | cur_input = cur_input[:max_len] |
| 468 | elif self.mode == 'mid' and len(cur_input) > max_len: |
| 469 | mid_cut_len = max_len // 2 |
| 470 | cur_input = cur_input[:mid_cut_len] + cur_input[-mid_cut_len:] |
| 471 | |
| 472 | if left_padding: |
| 473 | # left padding with bos |
| 474 | tokens[i] = [self.tokenizer.bos_id] * padding_len + cur_input |
| 475 | else: |
| 476 | tokens[i] = cur_input + [self.pad_id] * padding_len |
| 477 | |
| 478 | return torch.LongTensor(tokens).cuda() |
| 479 | |
| 480 | def batch_decode(self, |
| 481 | outputs, |
no outgoing calls
no test coverage detected