| """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) |
|
|