| 5 | |
| 6 | |
| 7 | class LMEvalAdaptor(BaseLM): |
| 8 | |
| 9 | def __init__(self, model_name, model, tokenizer, batch_size=1, max_length=-1): |
| 10 | super().__init__() |
| 11 | |
| 12 | assert isinstance(batch_size, int) |
| 13 | |
| 14 | self.model_name = model_name |
| 15 | self.model = model |
| 16 | self.model.eval() |
| 17 | |
| 18 | self.tokenizer = tokenizer |
| 19 | |
| 20 | # assert isinstance(self.tokenizer, ( |
| 21 | # transformers.GPT2Tokenizer, transformers.GPT2TokenizerFast, |
| 22 | # transformers.T5Tokenizer, transformers.T5TokenizerFast, |
| 23 | # )), "this tokenizer has not been checked for compatibility yet!" |
| 24 | |
| 25 | self.vocab_size = self.tokenizer.vocab_size |
| 26 | |
| 27 | self._batch_size = batch_size |
| 28 | |
| 29 | self._max_length = max_length |
| 30 | |
| 31 | @property |
| 32 | def eot_token_id(self): |
| 33 | # we use EOT because end of *text* is more accurate for what we're doing than end of *sentence* |
| 34 | return self.tokenizer.eos_token_id |
| 35 | |
| 36 | @property |
| 37 | def max_length(self): |
| 38 | if self._max_length != -1: |
| 39 | return self._max_length |
| 40 | if hasattr(self.model.config, 'n_ctx'): |
| 41 | return self.model.config.n_ctx |
| 42 | elif hasattr(self.model.config, 'max_position_embeddings'): |
| 43 | return self.model.config.max_position_embeddings |
| 44 | elif hasattr(self.model.config, 'n_positions'): |
| 45 | return self.model.config.n_positions |
| 46 | elif 'bloom' in self.model_name: |
| 47 | return 2048 |
| 48 | elif 'llama' in self.model_name: |
| 49 | return 2048 # TODO: did not check this |
| 50 | elif 'mpt' in self.model_name: |
| 51 | return 2048 |
| 52 | elif 'falcon' in self.model_name: |
| 53 | return 2048 |
| 54 | else: |
| 55 | print(self.model.config) |
| 56 | raise NotImplementedError |
| 57 | |
| 58 | @property |
| 59 | def max_gen_toks(self): |
| 60 | return 256 |
| 61 | |
| 62 | @property |
| 63 | def batch_size(self): |
| 64 | return self._batch_size |