File size: 7,833 Bytes
3f98d52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
"""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()