"""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()