ReMEDi / scripts /run_ablations.py
pranamanam's picture
Upload 62 files
3f98d52 verified
Raw
History Blame Contribute Delete
1.3 kB
"""Run uncertainty ablations with variant-specific calibration."""
import argparse
from pathlib import Path
import joblib
from remedi.io import ResponseData
from remedi.pipeline import calibrate,evaluate
from remedi.plotting import plot_summary
p=argparse.ArgumentParser()
for name in ['data','model','output']:p.add_argument('--'+name,required=True)
p.add_argument('--tolerance',type=float,required=True)
p.add_argument('--max-queries',type=int,default=20)
p.add_argument('--scenarios',type=int,default=32)
p.add_argument('--seed',type=int,default=0)
p.add_argument('--embeddings')
a=p.parse_args();d=ResponseData.load(a.data);m=joblib.load(Path(a.model)/'model.joblib')
for variant in ['none','permuted','diagonal','no-floor','no-target-sampling']:
out=Path(a.output)/variant
cal_variant='none' if variant=='no-target-sampling' else variant
c=calibrate(d,m,out/'calibration',ablation=cal_variant,scenarios=a.scenarios,
seed=a.seed,embeddings=a.embeddings,panel_molecules=3)
evaluate(d,m,c,out/'evaluation',ablation=variant,scenarios=a.scenarios,
seed=a.seed,embeddings=a.embeddings,tolerance=a.tolerance,max_queries=a.max_queries)
plot_summary(out/'evaluation/summary.csv',out/'plots',title=variant)
print(f'Completed {variant}',flush=True)