PIVOT / src /pivot /data /perturb_data.py
pranamanam's picture
Upload 176 files
6fa9282 verified
Raw
History Blame Contribute Delete
5.18 kB
"""Expression cache with explicit partitions and a metadata-only gene vocabulary."""
from __future__ import annotations
import json
from pathlib import Path
import numpy as np
import pandas as pd
import scipy.sparse as sp
from .preprocess import parse
class PerturbData:
"""Load a prepared cache. Arrays retain the row order of obs.parquet."""
def __init__(self, cache_dir: str, embedding="pca"):
self.dir = str(cache_dir)
p = Path(cache_dir)
self.meta = json.loads((p / "meta.json").read_text())
if self.meta.get("protocol") != "split-first-v1":
raise ValueError(
"Use a split-first-v1 cache. Historical checkpoints require the archived code."
)
if embedding != "pca":
raise ValueError(
"This release implements the manuscript PCA representation"
)
self.obs = pd.read_parquet(p / "obs.parquet")
self.genes = (p / "genes_hvg.txt").read_text().splitlines()
self.Xhvg = sp.load_npz(p / "Xhvg.npz").tocsr()
self.emb = np.load(p / "pca_emb.npy")
self.d = self.emb.shape[1]
self.pca_components = np.load(p / "pca_components.npy")
self.pca_mean = np.load(p / "pca_mean.npy")
if len(self.obs) != len(self.emb) or self.Xhvg.shape != (
len(self.obs),
len(self.genes),
):
raise ValueError("Cache row or feature dimensions disagree")
self.sep = self.meta["sep"]
self.control_label = self.meta["control_label"]
self.operation = self.meta["operation"]
self.is_control = self.obs.is_control.to_numpy()
self.control_idx = np.flatnonzero(self.is_control)
self.batch = self.obs.batch.to_numpy()
self.celltype = self.obs.celltype.to_numpy()
self.pert_to_idx = {
p: s.index.to_numpy()
for p, s in self.obs.groupby("perturbation")
if p != self.control_label
}
self.perturbations = sorted(self.pert_to_idx)
self.genes_vocab = sorted(
{g for p in self.perturbations for g in self.parse(p)}
)
self.gene_to_id = {g: i for i, g in enumerate(self.genes_vocab)}
self.op_vocab = ["none", self.operation]
self.op_to_id = {o: i for i, o in enumerate(self.op_vocab)}
self.singles = [p for p in self.perturbations if len(self.parse(p)) == 1]
self.combos = [p for p in self.perturbations if len(self.parse(p)) == 2]
def parse(self, label: str) -> list[str]:
return parse(label, self.control_label, self.sep)
def indices(self, partition: str, controls: bool | None = None) -> np.ndarray:
mask = self.obs.split.eq(partition).to_numpy(copy=True)
if controls is not None:
mask &= self.is_control if controls else ~self.is_control
return np.flatnonzero(mask)
def labels(self, partition: str) -> list[str]:
return sorted(set(self.obs.iloc[self.indices(partition, False)].perturbation))
def decode_to_genes(self, emb: np.ndarray) -> np.ndarray:
"""Rank-d affine reconstruction, shape (..., d) to (..., n_hvg)."""
return emb @ self.pca_components + self.pca_mean
def pert_gene_op_ids(self, label):
genes = self.parse(label)
unknown = set(genes) - set(self.gene_to_id)
if unknown:
raise ValueError(f"Genes outside checkpoint vocabulary: {sorted(unknown)}")
ids = np.array([self.gene_to_id[g] for g in genes], dtype=np.int64)
return ids, np.full(len(ids), 1, dtype=np.int64)
def sample_controls(self, target_idx, strategy, rng, control_pool):
"""Sample only supplied controls, optionally within batch/cell type.
Missing matched controls raise an error so cross-context substitutions
cannot change the conditioning distribution without an explicit choice.
"""
pool = np.asarray(control_pool, dtype=int)
if not len(pool) or not self.is_control[pool].all():
raise ValueError("Invalid control pool")
if strategy == "random":
return rng.choice(pool, len(target_idx), replace=True)
if strategy == "nearest":
from sklearn.neighbors import NearestNeighbors
j = (
NearestNeighbors(n_neighbors=1)
.fit(self.emb[pool])
.kneighbors(self.emb[target_idx], return_distance=False)
)
return pool[j.ravel()]
if strategy not in ("batch", "celltype", "batch_celltype"):
raise ValueError(strategy)
out = []
for i in target_idx:
m = np.ones(len(pool), bool)
if strategy in ("batch", "batch_celltype"):
m &= self.batch[pool] == self.batch[i]
if strategy in ("celltype", "batch_celltype"):
m &= self.celltype[pool] == self.celltype[i]
if not m.any():
raise ValueError(
f"No matching control for cell {self.obs.iloc[i].cell_id}"
)
out.append(rng.choice(pool[m]))
return np.asarray(out, dtype=np.int64)