Download evaluate_glint.py from SlayerLab/Slayer149: direct link, hf CLI and curl.
- Browser
- Download file 2.59 kB
-
https://huggingface.co/SlayerLab/Slayer149/resolve/main/evaluate_glint.py
- Command line
-
hf download hf://SlayerLab/Slayer149/evaluate_glint.py
-
curl -L -o evaluate_glint.py https://huggingface.co/SlayerLab/Slayer149/resolve/main/evaluate_glint.py
2.59 kB
| """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 | |
| 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() | |