| """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) |
| |
| 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.') |
| |
| 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()) |
|
|