Slayer149 / evaluate_glint.py
kacperwikiel's picture
Release GoLLeM 149M after 20B continuation tokens with full GLINT evaluation
3f431df verified
Raw History Blame Contribute Delete
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
@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()