File size: 5,182 Bytes
6fa9282
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
"""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)