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