File size: 1,297 Bytes
3f98d52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
"""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)