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