File size: 2,396 Bytes
80aea5b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
"""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()