Download scripts/run_ablations.py from ChatterjeeLab/ReMEDi: direct link, hf CLI and curl.
- Browser
- Download file 1.3 kB
-
https://huggingface.co/ChatterjeeLab/ReMEDi/resolve/main/scripts/run_ablations.py
- Command line
-
hf download hf://ChatterjeeLab/ReMEDi/scripts/run_ablations.py
-
curl -L -o run_ablations.py https://huggingface.co/ChatterjeeLab/ReMEDi/resolve/main/scripts/run_ablations.py
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) | |