ReMEDi / scripts /validate_release.py
pranamanam's picture
Upload 62 files
3f98d52 verified
Raw
History Blame Contribute Delete
2.24 kB
"""Reproduce the bundled measured-data checks, including frozen external transfer."""
import argparse
import json
from pathlib import Path
import subprocess
import sys
import joblib
import pandas as pd
import yaml
from remedi.data import prepare_h5ad
from remedi.pipeline import evaluate
from remedi.plotting import plot_summary
p=argparse.ArgumentParser()
p.add_argument('--output',default='runs/validation')
a=p.parse_args()
root=Path(__file__).resolve().parents[1]
out=Path(a.output).resolve();out.mkdir(parents=True,exist_ok=True)
def run(*args):
subprocess.run([sys.executable,*map(str,args)],check=True,cwd=root)
for study,filename in [('tahoe','tahoe.h5ad'),('sciplex','sciplex_prototype.h5ad')]:
run(root/'scripts/run_pipeline.py','--h5ad',root/'examples'/study/filename,
'--mapping',root/'configs'/f'{study}.yaml','--structures',root/'examples'/study/'structures.csv',
'--output',out/study,'--min-cells',2,'--dimensions',8,'--genes',250,'--max-queries',4,'--tolerance',1)
run(root/'scripts/align_genes.py','--inputs',root/'examples/tahoe/tahoe.h5ad',
root/'examples/sciplex/sciplex_prototype.h5ad','--gene-columns','gene_symbol','index',
'--drop-ambiguous','--output',out/'aligned')
run(root/'scripts/run_pipeline.py','--h5ad',out/'aligned/0_tahoe.h5ad',
'--mapping',root/'configs/tahoe.yaml','--structures',root/'examples/tahoe/structures.csv',
'--output',out/'transfer-source','--min-cells',2,'--dimensions',8,'--genes',250,'--max-queries',4,'--tolerance',1)
s=pd.read_csv(root/'examples/sciplex/structures.csv')
splits=s[['smiles']].copy();splits['split']='test'
d=prepare_h5ad(out/'aligned/1_sciplex_prototype.h5ad',out/'transfer/data',
yaml.safe_load((root/'configs/sciplex.yaml').read_text()),splits,s,
feature_model=out/'transfer-source/data/cell_feature_model.joblib',min_cells=2,max_cells=32)
m=joblib.load(out/'transfer-source/model/model.joblib')
c=json.loads((out/'transfer-source/calibration/calibration.json').read_text())
evaluate(d,m,c,out/'transfer/evaluation',external_track='unseen',scenarios=16,max_queries=4,tolerance=1)
plot_summary(out/'transfer/evaluation/summary.csv',out/'transfer/plots')
print(f'Completed measured-data integration checks in {out}')