Download src/pivot/cli.py from ChatterjeeLab/PIVOT: direct link, hf CLI and curl.
- Browser
- Download file 9.22 kB
-
https://huggingface.co/ChatterjeeLab/PIVOT/resolve/main/src/pivot/cli.py
- Command line
-
hf download hf://ChatterjeeLab/PIVOT/src/pivot/cli.py
-
curl -L -o cli.py https://huggingface.co/ChatterjeeLab/PIVOT/resolve/main/src/pivot/cli.py
9.22 kB
| """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() | |