MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / LMEvalAdaptor

Class LMEvalAdaptor

test/general/utils_eval.py:7–113  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

5
6
7class 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

Callers 1

llm_eval.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected