PIVOT / src /pivot /cli.py
pranamanam's picture
Upload 176 files
6fa9282 verified
Raw
History Blame Contribute Delete
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()