CoolFace
Modelpublic

1bitLLM/bitnet_b1_58-large

sourceHugging Facemitupdated 2y agoView on Hugging Face
129likes2.7kdownloads
eval_utils.py133 linesDownload Raw Back to root
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        )