| """Command-line data preparation, fitting, prediction, nomination and evaluation.""" |
|
|
| from __future__ import annotations |
| import argparse, json |
| from pathlib import Path |
| import numpy as np |
| import torch |
| from pivot.data.preprocess import prepare |
| from pivot.data.perturb_data import PerturbData |
| from pivot.training.train import TrainConfig, train, load_checkpoint |
| from pivot.evaluation.inference import ( |
| encode_label, |
| forward_predict, |
| endpoint_ranking, |
| reward_guidance, |
| project_and_rerank, |
| greedy_combinatorial, |
| ) |
| from pivot.evaluation.rewards import Reward |
| from pivot.evaluation.runner import evaluate, save_json |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser( |
| prog="pivot", |
| description="Transcriptomic endpoint prediction and intervention nomination", |
| ) |
| sub = parser.add_subparsers(dest="command", required=True) |
| p = sub.add_parser("prepare") |
| p.add_argument("--raw", required=True) |
| p.add_argument("--output", required=True) |
| p.add_argument("--dataset", choices=["norman", "replogle_k562"], required=True) |
| p.add_argument( |
| "--split", |
| dest="regime", |
| choices=["cell", "perturbation", "combination", "gene"], |
| default="perturbation", |
| ) |
| p.add_argument("--seed", type=int, default=0) |
| p.add_argument("--input-scale", choices=["counts", "log1p"], default="counts") |
| for key, default in ( |
| ("n-hvg", 2000), |
| ("n-pca", 50), |
| ("min-cells", 20), |
| ("max-cells", None), |
| ("max-per-group", None), |
| ("max-groups", None), |
| ): |
| p.add_argument("--" + key, type=int, default=default) |
| p.add_argument("--batch-col", default="gemgroup") |
| p.add_argument("--pert-col", default="perturbation") |
| p.add_argument("--celltype-col", default="celltype") |
| p = sub.add_parser("train") |
| p.add_argument("--cache", required=True) |
| p.add_argument("--config", required=True) |
| p.add_argument("--output", required=True) |
| p.add_argument("--resume") |
| p.add_argument("--device") |
| p.add_argument("--seed", type=int) |
| p = sub.add_parser("evaluate") |
| p.add_argument("--cache", required=True) |
| p.add_argument("--checkpoint") |
| p.add_argument("--output", required=True) |
| p.add_argument( |
| "--baseline", |
| choices=[ |
| "mean_control", |
| "average_effect", |
| "additive", |
| "ridge", |
| "endpoint_mlp", |
| "conditional_mlp", |
| ], |
| ) |
| p.add_argument("--partition", choices=["val", "test"], default="test") |
| p.add_argument( |
| "--catalog", choices=["single", "combination", "all"], default="single" |
| ) |
| p.add_argument( |
| "--reward", |
| dest="reward_kind", |
| choices=["cosine", "centroid", "mmd", "wasserstein"], |
| default="cosine", |
| ) |
| p.add_argument("--n-cells", type=int, default=128) |
| p.add_argument("--seed", type=int, default=0) |
| p.add_argument("--device", default="cpu") |
| p.add_argument("--guidance-steps", type=int, default=25) |
| p.add_argument("--step-size", type=float, default=0.5) |
| p.add_argument("--k-nearest", type=int, default=10) |
| p.add_argument("--initialization", choices=["best", "random"], default="best") |
| p.add_argument("--max-targets", type=int) |
| p.add_argument("--baseline-epochs", type=int, default=60) |
| p.add_argument("--baseline-hidden", type=int, default=512) |
| p.add_argument("--ridge-alpha", type=float, default=1.0) |
| p = sub.add_parser("predict") |
| p.add_argument("--cache", required=True) |
| p.add_argument("--checkpoint", required=True) |
| p.add_argument("--label", required=True) |
| p.add_argument("--output", required=True) |
| p.add_argument("--n-cells", type=int, default=128) |
| p.add_argument("--device", default="cpu") |
| p = sub.add_parser("nominate") |
| p.add_argument("--cache", required=True) |
| p.add_argument("--checkpoint", required=True) |
| p.add_argument( |
| "--target", |
| required=True, |
| help="Numpy (n,d) population in this cache's PCA basis", |
| ) |
| p.add_argument("--output", required=True) |
| p.add_argument( |
| "--catalog", choices=["single", "combination", "all"], default="single" |
| ) |
| p.add_argument( |
| "--reward", |
| choices=["centroid", "cosine", "mmd", "wasserstein"], |
| default="centroid", |
| ) |
| p.add_argument( |
| "--search", choices=["exhaustive", "guidance", "greedy"], default="exhaustive" |
| ) |
| p.add_argument("--steps", type=int, default=25) |
| p.add_argument("--max-size", type=int, default=2) |
| p.add_argument("--device", default="cpu") |
| p.add_argument("--n-cells", type=int, default=128) |
| p.add_argument("--seed", type=int, default=0) |
| a = vars(parser.parse_args()) |
| command = a.pop("command") |
| if command == "prepare": |
| print(json.dumps(prepare(**a), indent=2)) |
| return |
| data = PerturbData(a.pop("cache")) |
| if command == "train": |
| cfg = json.loads(Path(a.pop("config")).read_text()) |
| for k in ("device", "seed"): |
| v = a.pop(k) |
| if v is not None: |
| cfg[k] = v |
| train(data, TrainConfig(**cfg), **a) |
| return |
| if command == "evaluate": |
| ck = a.pop("checkpoint") |
| baseline = a.pop("baseline") |
| if bool(ck) == bool(baseline): |
| parser.error("Specify one of --checkpoint or --baseline") |
| epochs = a.pop("baseline_epochs") |
| hidden = a.pop("baseline_hidden") |
| alpha = a.pop("ridge_alpha") |
| if baseline: |
| from pivot.evaluation.baselines import Baseline |
|
|
| b = Baseline(data, baseline, alpha, epochs, hidden, a["seed"], a["device"]) |
| result = evaluate(data, b.predict, method=baseline, **a) |
| result["baseline_fit"] = b.training_info |
| save_json(a["output"], result) |
| else: |
| model, cfg = load_checkpoint(ck, data, a["device"]) |
|
|
| def predict(c0, label): |
| return ( |
| forward_predict( |
| model, |
| torch.as_tensor(c0, device=a["device"]), |
| encode_label(model, data, label, a["device"]), |
| ) |
| .cpu() |
| .numpy() |
| ) |
|
|
| result = evaluate(data, predict, model=model, **a) |
| print(json.dumps(result["summary"], indent=2)) |
| return |
| model, cfg = load_checkpoint(a["checkpoint"], data, a["device"]) |
| rng = np.random.default_rng(a.get("seed", 0)) |
| ids = data.indices("test", True) |
| ids = rng.choice(ids, min(a["n_cells"], len(ids)), replace=False) |
| c0 = torch.as_tensor(data.emb[ids], device=a["device"]) |
| if command == "predict": |
| pred = ( |
| forward_predict( |
| model, c0, encode_label(model, data, a["label"], a["device"]) |
| ) |
| .cpu() |
| .numpy() |
| ) |
| Path(a["output"]).parent.mkdir(parents=True, exist_ok=True) |
| np.savez_compressed( |
| a["output"], |
| latent=pred, |
| expression_reconstruction=data.decode_to_genes(pred), |
| genes=np.asarray(data.genes), |
| control_cell_ids=data.obs.iloc[ids].cell_id.to_numpy(dtype=str), |
| ) |
| return |
| target = np.load(a["target"]) |
| if target.ndim != 2 or target.shape[1] != data.d or not np.isfinite(target).all(): |
| raise ValueError("Target needs finite (n,d) PCA coordinates") |
| reward = Reward( |
| a["reward"], |
| target_sample=target, |
| control_ref=c0.mean(0), |
| gamma=data.meta["mmd_gamma"], |
| device=a["device"], |
| ) |
| labels = ( |
| data.singles |
| if a["catalog"] == "single" |
| else data.combos if a["catalog"] == "combination" else data.perturbations |
| ) |
| if a["search"] == "greedy": |
| genes, score, history = greedy_combinatorial( |
| model, data, data.genes_vocab, c0, reward, a["max_size"], device=a["device"] |
| ) |
| out = { |
| "genes": genes, |
| "score": score, |
| "history": history, |
| "catalog_membership": data.sep.join(sorted(genes)) in labels, |
| } |
| else: |
| ranked = endpoint_ranking(model, data, labels, c0, reward, device=a["device"]) |
| if a["search"] == "guidance": |
| es = reward_guidance( |
| model, |
| c0, |
| reward, |
| encode_label(model, data, ranked[0][0], a["device"]), |
| a["steps"], |
| ) |
| ranked = project_and_rerank( |
| model, data, labels, es, c0, reward, device=a["device"] |
| ) |
| out = { |
| "ranked": [ |
| { |
| "label": q, |
| "genes": data.parse(q), |
| "operation": data.operation, |
| "predicted_reward": v, |
| } |
| for q, v in ranked |
| ] |
| } |
| out.update( |
| search=a["search"], |
| reward=a["reward"], |
| target_file=Path(a["target"]).name, |
| control_cell_ids=data.obs.iloc[ids].cell_id.tolist(), |
| data_meta=data.meta, |
| ) |
| save_json(a["output"], out) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|