"""Reproducible command-line entry points for the complete reference pipeline.""" import argparse import json from pathlib import Path import joblib import pandas as pd import yaml from .io import ResponseData def parser(): p = argparse.ArgumentParser(prog="remedi", description="Molecule-dose endpoint generation") sub = p.add_subparsers(dest="command", required=True) d = sub.add_parser("demo", help="Run a synthetic end-to-end engineering check") d.add_argument("--output", required=True); d.add_argument("--seed", type=int, default=0) d.add_argument("--max-queries", type=int, default=8) s = sub.add_parser("split", help="Create grouped structure/scaffold assignments") s.add_argument("--molecules", required=True); s.add_argument("--output", required=True) s.add_argument("--seed", type=int, default=0); s.add_argument("--scaffold", action="store_true") s.add_argument("--fractions", type=float, nargs=4, default=[.6,.15,.1,.15]) q = sub.add_parser("prepare", help="Aggregate a local RNA/embedding AnnData cohort") q.add_argument("--h5ad", required=True); q.add_argument("--mapping", required=True) q.add_argument("--splits", required=True); q.add_argument("--structures") q.add_argument("--output", required=True); q.add_argument("--representation", default="pca") q.add_argument("--feature-model"); q.add_argument("--max-cells", type=int, default=128) q.add_argument("--min-cells", type=int, default=8); q.add_argument("--genes", type=int, default=1000) q.add_argument("--dimensions", type=int, default=32); q.add_argument("--seed", type=int, default=0) q.add_argument("--layer") t = sub.add_parser("train", help="Fit grouped response ensemble and cross-fitted residuals") t.add_argument("--data", required=True); t.add_argument("--output", required=True) t.add_argument("--head", choices=["ridge","mlp"], default="ridge") t.add_argument("--alpha", type=float, default=10.); t.add_argument("--hidden", type=int, default=128) t.add_argument("--ensemble", type=int, default=5); t.add_argument("--folds", type=int, default=3) t.add_argument("--encoder", choices=["fingerprint","cached"], default="fingerprint") t.add_argument("--embeddings"); t.add_argument("--seed", type=int, default=0) c = sub.add_parser("calibrate", help="Estimate an empirical held-out panel radius") for name in ["data","model","output"]: c.add_argument("--"+name, required=True) c.add_argument("--coverage", type=float, default=.9); c.add_argument("--panel-molecules", type=int, default=6) c.add_argument("--scenarios", type=int, default=32); c.add_argument("--seed", type=int, default=0) c.add_argument("--ablation", choices=["none","permuted","diagonal","no-floor"], default="none") c.add_argument("--embeddings") e = sub.add_parser("evaluate", help="Compare generators on hidden measured candidate responses") for name in ["data","model","calibration","output"]: e.add_argument("--"+name, required=True) e.add_argument("--split", default="test"); e.add_argument("--scenarios", type=int, default=32) e.add_argument("--tau", type=float, default=.05); e.add_argument("--tolerance", type=float, required=True, help="Endpoint success threshold fixed on validation data") e.add_argument("--seed", type=int, default=0); e.add_argument("--max-queries", type=int, default=30) e.add_argument("--methods", nargs="+"); e.add_argument("--embeddings"); e.add_argument("--metric-data") e.add_argument("--external-track", choices=["none","shared","unseen"], default="none") e.add_argument("--ablation", choices=["none","permuted","diagonal","no-floor","no-target-sampling"], default="none") g = sub.add_parser("generate", help="Sample exact members of a supplied verified chemical catalog") for name in ["data","model","calibration","catalog","output"]: g.add_argument("--"+name, required=True) g.add_argument("--target-rows", type=int, nargs="+", required=True) g.add_argument("--doses", type=float, nargs="+", default=[.05,.5,5.]) g.add_argument("--scenarios", type=int, default=32); g.add_argument("--tau", type=float, default=.05) g.add_argument("--shortlist", type=int, default=512); g.add_argument("--draws", type=int, default=20) g.add_argument("--seed", type=int, default=0); g.add_argument("--embeddings") z = sub.add_parser("plot", help="Draw Ubuntu plots from measured result CSVs") z.add_argument("--summary", required=True); z.add_argument("--output", required=True) z.add_argument("--font-dir"); z.add_argument("--title", default="Retrospective endpoint evaluation") r = sub.add_parser("plot-curve", help="Plot calibration, scaling, or runtime records") for name in ["csv","output","x","y"]: r.add_argument("--"+name, required=True) r.add_argument("--hue", default="method"); r.add_argument("--lower"); r.add_argument("--upper"); r.add_argument("--font-dir") m = sub.add_parser("encode-molecules", help="Cache a revision-pinned frozen MolFormer checkpoint") m.add_argument("--molecules", required=True); m.add_argument("--output", required=True) m.add_argument("--model-id", default="ibm/MoLFormer-XL-both-10pct"); m.add_argument("--revision", required=True) m.add_argument("--batch-size", type=int, default=32); m.add_argument("--device", default="cpu") return p def main(argv=None): a = vars(parser().parse_args(argv)); command = a.pop("command") if command == "demo": from .demo import run_demo print(run_demo(**a)); return if command == "split": from .splits import make_splits frame = make_splits(pd.read_csv(a.pop("molecules")), seed=a["seed"], fractions=a["fractions"], scaffold=a["scaffold"]) Path(a["output"]).parent.mkdir(parents=True, exist_ok=True); frame.to_csv(a["output"], index=False); return if command == "prepare": from .data import prepare_h5ad a["mapping"] = yaml.safe_load(Path(a["mapping"]).read_text()) a["splits"] = pd.read_csv(a["splits"]) a["structures"] = pd.read_csv(a["structures"]) if a["structures"] else None a["path"] = a.pop("h5ad"); data = prepare_h5ad(**a) print(f"Prepared {len(data.obs)} conditions in {data.response.shape[1]} dimensions"); return if command == "train": from .models import train a["data"] = ResponseData.load(a["data"]); a["kind"] = a.pop("head") train(**a); print("Saved fitted ensemble and cross-fitted residuals"); return if command in ["calibrate","evaluate","generate"]: from .pipeline import calibrate, evaluate, generate_catalog a["data"] = ResponseData.load(a["data"]) a["model"] = joblib.load(Path(a["model"])/"model.joblib") if "calibration" in a: a["calibration"] = json.loads((Path(a["calibration"])/"calibration.json").read_text()) if command == "calibrate": print(json.dumps(calibrate(**a), indent=2)); return if command == "evaluate": if a["metric_data"]: a["metric_data"] = ResponseData.load(a["metric_data"]) frame = evaluate(**a); print(f"Evaluated {frame.query_id.nunique()} queries"); return a["catalog"] = pd.read_csv(a["catalog"]); a["target_indices"] = a.pop("target_rows") frame = generate_catalog(**a); print(f"Generated a distribution on {len(frame)} verified pairs"); return if command == "plot": from .plotting import plot_summary plot_summary(**a); return if command == "plot-curve": from .plotting import plot_curve plot_curve(**a); return if command == "encode-molecules": from .chemistry import cache_molformer a["smiles"] = pd.read_csv(a.pop("molecules")).smiles cache_molformer(**a) if __name__ == "__main__": main()