| """Held-out next-token loss/perplexity evaluation for a checkpoint and local corpus.""" | |
| import argparse, math, torch | |
| from .runtime import load_model | |
| from .tokenizer import ByteBPETokenizer | |
| def main(): | |
| p=argparse.ArgumentParser();p.add_argument('--checkpoint',required=True);p.add_argument('--tokenizer',required=True);p.add_argument('--text',required=True);a=p.parse_args() | |
| x=load_model(a.checkpoint,a.tokenizer);ids=x.tokenizer.encode(open(a.text,encoding='utf8').read(),True,True);n=x.model.config.max_seq_len;losses=[] | |
| with torch.no_grad(): | |
| for i in range(0,max(0,len(ids)-n),n): | |
| batch=torch.tensor([ids[i:i+n]],device=x.device) | |
| if batch.size(1)>2:_,loss,_=x.model(batch[:,:-1],batch);losses.append(loss.item()) | |
| loss=sum(losses)/len(losses);print({'mean_loss':loss,'perplexity':math.exp(loss),'windows':len(losses)}) | |
| if __name__=='__main__':main() | |