| 33 | |
| 34 | |
| 35 | class HuggingfaceTokenizer: |
| 36 | |
| 37 | def __init__(self, name, seq_len=None, clean=None, **kwargs): |
| 38 | assert clean in (None, 'whitespace', 'lower', 'canonicalize') |
| 39 | self.name = name |
| 40 | self.seq_len = seq_len |
| 41 | self.clean = clean |
| 42 | |
| 43 | # init tokenizer |
| 44 | self.tokenizer = AutoTokenizer.from_pretrained(name, **kwargs) |
| 45 | self.vocab_size = self.tokenizer.vocab_size |
| 46 | |
| 47 | def __call__(self, sequence, **kwargs): |
| 48 | return_mask = kwargs.pop('return_mask', False) |
| 49 | |
| 50 | # arguments |
| 51 | _kwargs = {'return_tensors': 'pt'} |
| 52 | if self.seq_len is not None: |
| 53 | _kwargs.update({ |
| 54 | 'padding': 'max_length', |
| 55 | 'truncation': True, |
| 56 | 'max_length': self.seq_len |
| 57 | }) |
| 58 | _kwargs.update(**kwargs) |
| 59 | |
| 60 | # tokenization |
| 61 | if isinstance(sequence, str): |
| 62 | sequence = [sequence] |
| 63 | if self.clean: |
| 64 | sequence = [self._clean(u) for u in sequence] |
| 65 | ids = self.tokenizer(sequence, **_kwargs) |
| 66 | |
| 67 | # output |
| 68 | if return_mask: |
| 69 | return ids.input_ids, ids.attention_mask |
| 70 | else: |
| 71 | return ids.input_ids |
| 72 | |
| 73 | def _clean(self, text): |
| 74 | if self.clean == 'whitespace': |
| 75 | text = whitespace_clean(basic_clean(text)) |
| 76 | elif self.clean == 'lower': |
| 77 | text = whitespace_clean(basic_clean(text)).lower() |
| 78 | elif self.clean == 'canonicalize': |
| 79 | text = canonicalize(basic_clean(text)) |
| 80 | return text |