d0rj's picture
Publish evaluated diffusion v2 with PLL intervals and TensorBoard traces
80aea5b verified
Raw History Blame Contribute Delete
2.4 kB
"""Reproduce the released model's eight-task likelihood evaluation."""
import argparse
import json
import importlib.metadata
from pathlib import Path
import torch
from transformers import AutoTokenizer,AutoModelForCausalLM,AutoModelForSeq2SeqLM,AutoModelForMaskedLM
from lm_eval import evaluator,tasks
from lm_eval.models.huggingface import HFLM
from adapters import UL2HFLM, DiffusionHFLM
def main():
if importlib.metadata.version('lm_eval') != '0.4.12':
raise RuntimeError('This reproduction protocol requires lm_eval==0.4.12')
p=argparse.ArgumentParser(description=__doc__)
p.add_argument('--device',default='cpu')
p.add_argument('--dtype',default='float32',choices=['float32','bfloat16'])
p.add_argument('--batch-size',type=int,default=1)
p.add_argument('--output',type=Path,required=True)
p.add_argument('--limit',type=int,help='Smoke only; not a full benchmark')
a=p.parse_args(); a.output.mkdir(parents=True,exist_ok=False)
torch.set_num_threads(4)
root=Path(__file__).resolve().parents[1]
config=json.loads((root/'config.json').read_text()); ul2=bool(config.get('ul2'))
cls=AutoModelForMaskedLM
if a.batch_size != 1: raise ValueError("Diffusion PLL requires --batch-size 1")
model=cls.from_pretrained(root,trust_remote_code=True,dtype=getattr(torch,a.dtype)).to(a.device).eval()
tok=AutoTokenizer.from_pretrained(root,trust_remote_code=True)
adapter=DiffusionHFLM(pretrained=model,tokenizer=tok,backend='seq2seq' if ul2 else 'causal',device=a.device,batch_size=a.batch_size,max_length=2048)
mapping={'hellaswag':'hellaswag','arc_easy':'arc','arc_challenge':'arc','piqa':'piqa','winogrande':'winogrande','openbookqa':'openbookqa','boolq':'super_glue/boolq','lambada_openai':'lambada'}
manager=tasks.TaskManager(include_defaults=False,include_path=sorted({Path(tasks.__file__).parent/v for v in mapping.values()}))
for name in mapping:
result=evaluator.simple_evaluate(model=adapter,tasks=[name],num_fewshot=0,limit=a.limit,bootstrap_iters=1000,log_samples=False,task_manager=manager,random_seed=1234,numpy_random_seed=1234,torch_random_seed=1234,fewshot_random_seed=1234,apply_chat_template=False)
def fallback(x): return x.item() if hasattr(x,'item') else str(x)
(a.output/(name+'.json')).write_text(json.dumps(result,indent=2,default=fallback))
if __name__=='__main__': main()