File size: 863 Bytes
6a251e9
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
"""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()