Download scripts/quickstart.py from ChatterjeeLab/PIVOT: direct link, hf CLI and curl.
- Browser
- Download file 3.28 kB
-
https://huggingface.co/ChatterjeeLab/PIVOT/resolve/main/scripts/quickstart.py
- Command line
-
hf download hf://ChatterjeeLab/PIVOT/scripts/quickstart.py
-
curl -L -o quickstart.py https://huggingface.co/ChatterjeeLab/PIVOT/resolve/main/scripts/quickstart.py
3.28 kB
| """Run the documented example from measured counts through nomination and plots.""" | |
| from pathlib import Path | |
| import argparse, subprocess, sys | |
| import numpy as np | |
| from pivot.data.perturb_data import PerturbData | |
| ROOT = Path(__file__).resolve().parents[1] | |
| def run(args): | |
| subprocess.run([sys.executable, "-m", "pivot.cli", *map(str, args)], check=True) | |
| if __name__ == "__main__": | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--output", default="runs/example") | |
| a = parser.parse_args() | |
| out = Path(a.output).resolve() | |
| if out.exists(): | |
| raise FileExistsError("Choose a new example output directory") | |
| cache = out / "cache" | |
| model = out / "model" | |
| run( | |
| [ | |
| "prepare", | |
| "--raw", | |
| ROOT / "fixtures/norman_small.h5ad", | |
| "--dataset", | |
| "norman", | |
| "--split", | |
| "perturbation", | |
| "--n-hvg", | |
| 200, | |
| "--n-pca", | |
| 10, | |
| "--output", | |
| cache, | |
| ] | |
| ) | |
| run( | |
| [ | |
| "train", | |
| "--cache", | |
| cache, | |
| "--config", | |
| ROOT / "configs/small.json", | |
| "--output", | |
| model, | |
| ] | |
| ) | |
| run( | |
| [ | |
| "evaluate", | |
| "--cache", | |
| cache, | |
| "--checkpoint", | |
| model / "best.pt", | |
| "--catalog", | |
| "all", | |
| "--n-cells", | |
| 16, | |
| "--guidance-steps", | |
| 3, | |
| "--output", | |
| out / "pivot.json", | |
| ] | |
| ) | |
| run( | |
| [ | |
| "evaluate", | |
| "--cache", | |
| cache, | |
| "--baseline", | |
| "ridge", | |
| "--catalog", | |
| "all", | |
| "--n-cells", | |
| 16, | |
| "--output", | |
| out / "ridge.json", | |
| ] | |
| ) | |
| data = PerturbData(str(cache)) | |
| label = data.labels("test")[0] | |
| ids = np.intersect1d(data.indices("test", False), data.pert_to_idx[label]) | |
| np.save(out / "target.npy", data.emb[ids]) | |
| run( | |
| [ | |
| "predict", | |
| "--cache", | |
| cache, | |
| "--checkpoint", | |
| model / "best.pt", | |
| "--label", | |
| label, | |
| "--n-cells", | |
| 16, | |
| "--output", | |
| out / "prediction.npz", | |
| ] | |
| ) | |
| for search in ["exhaustive", "guidance", "greedy"]: | |
| run( | |
| [ | |
| "nominate", | |
| "--cache", | |
| cache, | |
| "--checkpoint", | |
| model / "best.pt", | |
| "--target", | |
| out / "target.npy", | |
| "--catalog", | |
| "all", | |
| "--search", | |
| search, | |
| "--steps", | |
| 3, | |
| "--n-cells", | |
| 16, | |
| "--output", | |
| out / (search + ".json"), | |
| ] | |
| ) | |
| subprocess.run( | |
| [ | |
| sys.executable, | |
| str(ROOT / "scripts/plot_results.py"), | |
| str(out / "pivot.json"), | |
| str(out / "ridge.json"), | |
| "--output", | |
| str(out / "plots"), | |
| ], | |
| check=True, | |
| ) | |
| print("Example complete:", out) | |