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