ganesh333's picture
Upload folder using huggingface_hub
575e16b verified
Raw History Blame Contribute Delete
1.23 kB
"""Re-score an existing CNN model dir on the official test split (and its session-clean subset)."""
import argparse, json
from pathlib import Path
from src.inference import load_model
from train import load_rows, evaluate_test, ROOT
p=argparse.ArgumentParser()
p.add_argument('--model-dir',default=str(ROOT)); p.add_argument('--data-dir',default=str(ROOT/'data'))
p.add_argument('--eval-crops',type=int,default=None,help='defaults to eval_crops in the model config (1 if absent)')
p.add_argument('--name',default=None,help='label for compare_models.py; result saved to eval/<name>.json')
a=p.parse_args()
mdir=Path(a.model_dir) if Path(a.model_dir).exists() else ROOT/a.model_dir
model,labels,cfg=load_model(mdir)
k=a.eval_crops if a.eval_crops is not None else int(cfg.get('eval_crops',1))
ev=evaluate_test(model,load_rows(a.data_dir),k,labels)
res={'model_dir':str(mdir),'eval_crops':k,**ev}
name=a.name or f"{mdir.resolve().name}_k{k}"
out=Path(__file__).resolve().parent/'eval'; out.mkdir(exist_ok=True)
(out/f'{name}.json').write_text(json.dumps(res,indent=2))
summary={s:{m:v for m,v in (ev[s] or {}).items() if m!='report'} for s in ev}
print(json.dumps({'name':name,'eval_crops':k,**summary},indent=2))