1bitLLM/bitnet_b1_58-large
1292.7k
1import torch2 3import numpy as np4import torch.nn.functional as F5 6from lm_eval.base import BaseLM7from datasets import load_dataset8 9 10def set_seed(seed):11 np.random.seed(seed)12 torch.random.manual_seed(seed)13 14def get_test_dataset(dataset_name, tokenizer, seqlen=2048):15 if dataset_name == "wikitext2":16 testdata = load_dataset('wikitext', 'wikitext-2-raw-v1', split='test')17 testdata = "".join(testdata['text']).split('\n')18 elif dataset_name == "c4":19 testdata = load_dataset('allenai/c4', data_files={'validation': 'en/c4-validation.00000-of-00008.json.gz'}, split='validation')['text']20 else:21 raise NotImplementedError22 23 testdata = [item for item in testdata if item != ""]24 tokenized_text = [tokenizer(item, add_special_tokens=False)['input_ids'] + [tokenizer.eos_token_id] for item in testdata]25 26 data, doc = [], [tokenizer.bos_token_id]27 for sen in tokenized_text:28 if len(sen) > seqlen:29 continue30 if len(doc) + len(sen) > seqlen:31 data.append(doc)32 doc = [tokenizer.bos_token_id]33 doc.extend(sen)34 if len(doc) > 1 and len(doc) <= seqlen:35 data.append(doc)36 return data37 38 39class LMEvalAdaptor(BaseLM):40 def __init__(self, model_name, model, tokenizer, batch_size=1, max_length=-1):41 super().__init__()42 43 assert isinstance(batch_size, int)44 45 self.model_name = model_name46 self.model = model47 self.model.eval()48 49 self.tokenizer = tokenizer50 51 self.vocab_size = self.tokenizer.vocab_size52 53 self._batch_size = batch_size54 55 self._max_length = max_length56 57 @property58 def eot_token_id(self):59 # we use EOT because end of *text* is more accurate for what we're doing than end of *sentence*60 return self.tokenizer.eos_token_id61 62 @property63 def max_length(self):64 if self._max_length != -1:65 return self._max_length66 if hasattr(self.model.config, "n_ctx"):67 return self.model.config.n_ctx68 elif hasattr(self.model.config, "max_position_embeddings"):69 return self.model.config.max_position_embeddings70 elif hasattr(self.model.config, "n_positions"):71 return self.model.config.n_positions72 elif "bloom" in self.model_name:73 return 204874 elif "llama" in self.model_name:75 return 2048 # TODO: did not check this76 elif "mpt" in self.model_name:77 return 204878 elif "falcon" in self.model_name:79 return 204880 else:81 print(self.model.config)82 raise NotImplementedError83 84 @property85 def max_gen_toks(self):86 return 25687 88 @property89 def batch_size(self):90 return self._batch_size91 92 @property93 def device(self):94 return "cuda"95 96 def tok_encode(self, string: str, add_special_tokens=True):97 return self.tokenizer.encode(string, add_special_tokens=add_special_tokens)98 99 def tok_decode(self, tokens):100 return self.tokenizer.decode(tokens)101 102 def loglikelihood(self, requests):103 new_reqs = []104 for context, continuation in requests:105 context, continuation = context.strip(), continuation.strip()106 if context == "":107 # end of text as context108 context_enc = [self.eot_token_id]109 else:110 context_enc = self.tok_encode(context, add_special_tokens=True)111 112 continuation_enc = self.tok_encode(continuation, add_special_tokens=False)113 114 new_reqs.append(((context, continuation), context_enc, continuation_enc))115 116 return self._loglikelihood_tokens(new_reqs)117 118 def _model_call(self, inps):119 """120 inps: a torch tensor of shape [batch, sequence]121 the size of sequence may vary from call to call122 123 returns: a torch tensor of shape [batch, sequence, vocab] with the124 logits returned from the model125 """126 with torch.no_grad():127 out = self.model(inps)[0]128 return out129 130 def _model_generate(self, context, max_length, eos_token_id):131 return self.model.generate(132 context, max_length=max_length, eos_token_id=eos_token_id, do_sample=False133 )