Ares-Lab / ares /evaluate.py
Ares Publisher
Add governed general-language 20M Colab experiment
6a251e9
Raw
History Blame Contribute Delete
863 Bytes
"""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()