ReMEDi / src /remedi /cli.py
pranamanam's picture
Upload 62 files
3f98d52 verified
Raw
History Blame Contribute Delete
7.83 kB
"""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()