ReMEDi / scripts /run_pipeline.py
pranamanam's picture
Upload 62 files
3f98d52 verified
Raw
History Blame Contribute Delete
3.86 kB
"""Run preparation, fitting, calibration, evaluation, generation, and figures."""
import argparse
from pathlib import Path
import json
import anndata as ad
import pandas as pd
import numpy as np
import yaml
from remedi.data import standardize_obs, prepare_h5ad
from remedi.splits import make_splits
from remedi.models import train
from remedi.pipeline import calibrate, evaluate, generate_catalog
from remedi.plotting import plot_summary
p=argparse.ArgumentParser()
p.add_argument('--h5ad',required=True)
p.add_argument('--mapping',required=True)
p.add_argument('--structures',required=True)
p.add_argument('--output',required=True)
p.add_argument('--seed',type=int,default=0)
p.add_argument('--max-queries',type=int,default=6)
p.add_argument('--min-cells',type=int,default=2)
p.add_argument('--max-cells',type=int,default=32)
p.add_argument('--dimensions',type=int,default=16)
p.add_argument('--genes',type=int,default=500)
p.add_argument('--ensemble',type=int,default=3)
p.add_argument('--scenarios',type=int,default=16)
p.add_argument('--tolerance',type=float,required=True,help='Preselected endpoint-success threshold')
p.add_argument('--representation',default='pca')
p.add_argument('--embeddings')
a=p.parse_args()
out=Path(a.output);out.mkdir(parents=True,exist_ok=True)
mapping=yaml.safe_load(Path(a.mapping).read_text())
structures=pd.read_csv(a.structures)
x=ad.read_h5ad(a.h5ad)
obs=standardize_obs(x.obs,mapping,structures)
# Retain compounds with an eligible two-block query at one or more doses.
counts=obs[~obs.is_control].groupby(['smiles','dose_um','context','block','control_id']).size()
valid=counts[counts>=a.min_cells].reset_index()
valid=valid.groupby(['smiles','dose_um','context']).block.nunique()
eligible=set(valid[valid>=2].reset_index().smiles)
keep=obs.is_control|obs.smiles.isin(eligible)
x=x[keep.to_numpy()].copy();x.write_h5ad(out/'cohort.h5ad',compression='gzip')
smiles=sorted(eligible)
if len(smiles)<9:raise ValueError(f'Need at least 9 compounds with two blocks. Found {len(smiles)}. Increase input cohort.')
# A small engineering cohort needs three calibration and two test compounds.
fractions=((len(smiles)-6)/len(smiles),1/len(smiles),3/len(smiles),2/len(smiles)) if len(smiles)<30 else (.55,.1,.2,.15)
splits=make_splits(pd.DataFrame({'smiles':smiles}),a.seed,fractions)
splits.to_csv(out/'splits.csv',index=False)
d=prepare_h5ad(out/'cohort.h5ad',out/'data',mapping,splits,structures,
representation=a.representation,max_cells=a.max_cells,min_cells=a.min_cells,
genes=a.genes,dimensions=a.dimensions,seed=a.seed)
m=train(d,out/'model',ensemble=a.ensemble,folds=3,seed=a.seed,
encoder='cached' if a.embeddings else 'fingerprint',embeddings=a.embeddings)
c=calibrate(d,m,out/'calibration',scenarios=a.scenarios,panel_molecules=3,seed=a.seed,embeddings=a.embeddings)
e=evaluate(d,m,c,out/'evaluation',scenarios=a.scenarios,max_queries=a.max_queries,
tolerance=a.tolerance,seed=a.seed,embeddings=a.embeddings)
plot_summary(out/'evaluation/summary.csv',out/'plots')
query=d.obs[d.obs.split=='test'].groupby(['molecule_id','dose_um','context'],sort=True)
indices=next(g.index.tolist() for _,g in query if len(g)>=2)
catalog=pd.DataFrame({'smiles':smiles,'catalog_id':[f'cohort_{i}' for i in range(len(smiles))]})
catalog.to_csv(out/'example_catalog.csv',index=False)
generate_catalog(d,m,c,catalog,indices,out/'generation',scenarios=a.scenarios,
seed=a.seed,embeddings=a.embeddings,doses=sorted(d.obs.dose_um.unique()))
(out/'run.json').write_text(json.dumps({'source':a.h5ad,'cells':x.n_obs,'genes':x.n_vars,
'compounds':len(smiles),'conditions':len(d.obs),'queries':int(e.query_id.nunique()),
'representation':a.representation,'seed':a.seed,'example_catalog':'measured cohort, not Enamine REAL'},indent=2)+'\n')
print((out/'run.json').read_text())