File size: 2,588 Bytes
3f431df
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
39
40
"""Full pinned GLINT-style evaluation of the released safetensors weights."""
import argparse,json,math
from pathlib import Path
import torch
import pyarrow.parquet as pq
from huggingface_hub import hf_hub_download,list_repo_files
from load_model import load_model
import glint_metrics as metrics

def main():
    p=argparse.ArgumentParser();p.add_argument('--model-dir',default='.');p.add_argument('--out',default='glint-results.json');a=p.parse_args()
    root=Path(a.model_dir);torch.set_num_threads(4)
    model,tok=load_model(root,'cuda')
    reference=json.loads((root/'evaluation.json').read_text())
    revisions=reference['dataset_revisions']
    files={r:list_repo_files(r,repo_type='dataset',revision=rev) for r,rev in revisions.items()}
    def rows(repo,config,split):
        fs=[f for f in files[repo] if f.endswith('.parquet') and f.split('/')[0]==config and (split in Path(f).name or '/'+split+'/' in '/'+f)]
        result=[]
        for f in sorted(fs):result.extend(pq.read_table(hf_hub_download(repo,f,repo_type='dataset',revision=revisions[repo])).to_pylist())
        if not result:raise ValueError(f'No data: {repo}/{config}/{split}')
        return result
    metrics._rows=rows
    @torch.inference_mode()
    def logits(x):
        with torch.autocast('cuda',dtype=torch.bfloat16):return model(x)[0].float()
    text=' '.join(r['text'] for r in rows('Salesforce/wikitext','wikitext-2-raw-v1','test')).strip()
    ppl=metrics.compute_perplexity(logits,tok,text,'cuda')
    byte_ppl=math.exp(math.log(ppl)*len(tok.encode(text).ids)/len(text.encode('utf-8')))
    result={'parameters':sum(p.numel() for p in model.parameters()),'dataset_revisions':revisions,'protocol':reference['protocol'],'wikitext2_token_ppl':ppl,'wikitext2_byte_ppl':byte_ppl}
    result.update(metrics.evaluate_blimp(logits,tok,'cuda'));result.update(metrics.evaluate_arc_easy(logits,tok,'cuda'))
    assert result['blimp_n']==67000 and result['arc_n']==2376
    board=json.loads((root/'board_snapshot.json').read_text())
    wiki=100*max(0,min(1,1-(math.log(min(byte_ppl,500))-board['wikiMinLog'])/(board['wikiMaxLog']-board['wikiMinLog'])))
    result['overall_score_fixed_board_snapshot']=(result['blimp_acc']+result['arc_easy_acc']+wiki)/3
    bonus=1+.5*max(0,min(1,(board['paramLogMax']-math.log10(result['parameters']))/(board['paramLogMax']-board['paramLogMin'])))
    result['efficiency_fixed_board_snapshot']=result['overall_score_fixed_board_snapshot']*bonus
    Path(a.out).write_text(json.dumps(result,indent=2));print(json.dumps(result,indent=2))
if __name__=='__main__':main()