File size: 9,770 Bytes
3f98d52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
"""AnnData ingestion with explicit units, control matching, and train-only features."""
from pathlib import Path
import json
import joblib
import numpy as np
import pandas as pd
from scipy import sparse
from sklearn.decomposition import PCA
from .chemistry import canonical_smiles, molecule_id
from .io import ResponseData, fingerprint_config

UNIT_TO_UM = {"uM": 1., "µM": 1., "μM": 1., "nM": .001, "mM": 1000., "M": 1e6}

def standardize_obs(obs, mapping, structures=None):
    """Map original fields without guessing dose units or replicate identities."""
    required = ["drug", "dose", "context", "block", "control_group"]
    for name in required:
        if name not in mapping or mapping[name] not in obs:
            raise ValueError(f"Mapping requires an existing {name} column")
    unit = mapping.get("dose_unit")
    if unit not in UNIT_TO_UM:
        raise ValueError("Specify dose_unit explicitly as uM, nM, mM, or M")
    result = pd.DataFrame(index=obs.index)
    for name in required:
        result[name] = obs[mapping[name]].to_numpy()
    result["is_control"] = result.drug.astype(str).isin(mapping.get("control_names", ["DMSO"]))
    result["dose_um"] = pd.to_numeric(result.dose, errors="raise")*UNIT_TO_UM[unit]
    result.loc[result.is_control, "dose_um"] = 0.
    for field in ["drug", "context", "block", "control_group"]:
        if result[field].isna().any(): raise ValueError(f"Missing {field} identifiers")
        result[field] = result[field].astype(str)
    result["control_id"] = result.context+"|"+result.control_group
    if "smiles" in mapping:
        result["smiles"] = obs[mapping["smiles"]].to_numpy()
    elif structures is not None:
        if structures.drug.duplicated().any(): raise ValueError("Ambiguous drug-to-structure mapping")
        result["smiles"] = result.drug.map(structures.set_index("drug").smiles)
    else:
        raise ValueError("Provide a SMILES column or a drug,smiles mapping file")
    result.loc[result.is_control, "smiles"] = "CS(C)=O"
    if result.smiles.isna().any():
        missing = result.loc[result.smiles.isna(), "drug"].unique().tolist()
        raise ValueError(f"Unmapped structures: {missing[:10]}")
    canonical = {s: canonical_smiles(s) for s in result.smiles.unique()}
    result["smiles"] = result.smiles.map(canonical)
    result["molecule_id"] = result.smiles.map(molecule_id)
    if not np.isfinite(result.dose_um).all(): raise ValueError("Concentrations must be finite")
    if ((result.dose_um <= 0) & ~result.is_control).any():
        raise ValueError("Non-control records have nonpositive concentration")
    result["source_cell_id"] = result.index.astype(str)
    return result.reset_index(drop=True)

def normalize_counts(x):
    if sparse.issparse(x):
        x = x.astype(np.float64).tocsr()
        if not np.isfinite(x.data).all() or (x.data < 0).any():
            raise ValueError("RNA input must contain finite nonnegative counts")
        total = np.asarray(x.sum(1)).ravel()
        x = sparse.diags(np.divide(1e4, total, out=np.zeros_like(total), where=total > 0))@x
        x.data = np.log1p(x.data)
        return x.tocsr()
    x = np.asarray(x)
    x = np.asarray(x, dtype=np.float64)
    if not np.isfinite(x).all() or (x < 0).any():
        raise ValueError("RNA input must contain finite nonnegative counts")
    total = x.sum(axis=1, keepdims=True)
    return np.log1p(np.divide(x*1e4, total, out=np.zeros_like(x), where=total > 0))

def prepare_h5ad(path, output, mapping, splits, structures=None, representation="pca",
                 feature_model=None, max_cells=128, min_cells=8, genes=1000, dimensions=32,
                 seed=0, layer=None, fit_cell_cap=4096):
    import anndata as ad
    a = ad.read_h5ad(path, backed="r")
    obs = standardize_obs(a.obs, mapping, structures)
    splits = splits.copy()
    splits["smiles"] = splits.smiles.map(canonical_smiles)
    if splits.smiles.duplicated().any(): raise ValueError("Split manifest has duplicate structures")
    obs["split"] = obs.smiles.map(splits.set_index("smiles").split)
    obs.loc[obs.is_control, "split"] = "control"
    if obs.loc[~obs.is_control, "split"].isna().any():
        raise ValueError("Split manifest is missing treated structures")
    rng = np.random.default_rng(seed)
    group_fields = ["molecule_id", "dose_um", "context", "block", "control_id", "is_control"]
    groups = []
    for _, group in obs.groupby(group_fields, sort=True):
        if len(group) < min_cells: continue
        idx = group.index.to_numpy()
        groups.append(np.sort(rng.choice(idx, min(max_cells, len(idx)), replace=False)))
    if not groups: raise ValueError("No groups pass the minimum-cell threshold")
    selected = np.sort(np.concatenate(groups))
    position = {int(v): k for k, v in enumerate(selected)}
    source = a.layers[layer] if layer else a.X
    if not a.var_names.is_unique: raise ValueError("Gene identifiers must be unique")
    if min_cells < 2 or max_cells < min_cells: raise ValueError("Require max_cells >= min_cells >= 2")
    if representation.startswith("obsm:"):
        name = representation.split(":", 1)[1]
        if name not in a.obsm: raise ValueError(f"Missing frozen cell embeddings in obsm[{name!r}]")
        features = np.asarray(a.obsm[name][selected], float)
        feature_metadata = {"representation": representation, "dimensions": features.shape[1],
                            "checkpoint": mapping.get("cell_checkpoint", "unrecorded")}
        if feature_metadata["checkpoint"] == "unrecorded":
            raise ValueError("Record cell_checkpoint in the mapping for frozen embeddings")
    elif representation == "pca":
        # Work in bounded batches. Feature selection sees training treated cells only.
        if feature_model:
            projector = joblib.load(feature_model)
            order = a.var_names.get_indexer(projector["genes"])
            if (order < 0).any():
                raise ValueError("External cohort misses fitted genes. Align both raw cohorts to a common gene set before fitting.")
        else:
            order = np.arange(a.n_vars)
            train = selected[obs.iloc[selected].split.to_numpy() == "train"]
            if len(train) < 3: raise ValueError("PCA requires training cells")
            train = np.sort(rng.choice(train, min(fit_cell_cap, len(train)), replace=False))
            x = normalize_counts(source[train])
            variance = np.asarray(x.power(2).mean(0)).ravel()-np.asarray(x.mean(0)).ravel()**2 if sparse.issparse(x) else x.var(0)
            keep = np.argsort(variance)[-min(genes, x.shape[1]):]
            d = min(dimensions, len(train)-1, len(keep))
            xx = x[:, keep].toarray() if sparse.issparse(x) else x[:, keep]
            pca = PCA(n_components=d, random_state=seed).fit(xx)
            projector = {"genes": list(a.var_names), "keep": keep, "pca": pca,
                         "seed": seed, "fit_cells": len(train)}
        pieces = []
        for start in range(0, len(selected), 512):
            x = normalize_counts(source[selected[start:start+512]][:, order])
            xx = x[:, projector["keep"]]
            pieces.append(projector["pca"].transform(xx.toarray() if sparse.issparse(xx) else xx))
        features = np.concatenate(pieces)
        Path(output).mkdir(parents=True, exist_ok=True)
        joblib.dump(projector, Path(output)/"cell_feature_model.joblib")
        feature_metadata = {"representation": "training-only log-count PCA", "dimensions": features.shape[1],
                            "feature_genes": [projector["genes"][i] for i in projector["keep"]],
                            "projection_hash": fingerprint_config({"components": projector["pca"].components_.tolist(), "mean": projector["pca"].mean_.tolist(), "normalization_genes": projector["genes"]})}
    else:
        raise ValueError("Cell representation must be pca or obsm:<frozen_embedding_key>")
    summaries = []
    controls = {}
    for idx in groups:
        row = obs.iloc[idx[0]].to_dict()
        f = features[[position[int(i)] for i in idx]]
        summary = (row, f.mean(0), f.var(0, ddof=1)/len(f), len(f))
        if row["is_control"]:
            # Multiple control samples in one pool are combined at the cell level.
            controls.setdefault(row["control_id"], []).append(f)
        else: summaries.append(summary)
    controls = {k: np.concatenate(v) for k, v in controls.items()}
    records, responses, contexts, tv, cv = [], [], [], [], []
    dropped = 0
    for row, mean, var, n in summaries:
        if row["control_id"] not in controls:
            dropped += 1; continue
        f = controls[row["control_id"]]
        c, v = f.mean(0), f.var(0, ddof=1)/len(f)
        row["n_cells"], row["n_control_cells"] = n, len(f)
        records.append({k: row[k] for k in ["molecule_id", "smiles", "drug", "dose_um", "context", "block", "control_id", "split", "n_cells", "n_control_cells"]})
        responses.append(mean-c); contexts.append(c); tv.append(var); cv.append(v)
    if a.isbacked: a.file.close()
    if not records: raise ValueError("No treated conditions have matching controls")
    metadata = {"source": str(path), "study": mapping.get("study"), "mapping": mapping, "seed": seed, "max_cells": max_cells,
                "min_cells": min_cells, "conditions_without_controls": dropped,
                "cell_features": feature_metadata, "feature_space_id": fingerprint_config(feature_metadata),
                "sampling_model": "block bootstrap plus diagonal Gaussian sampling error of population means"}
    data = ResponseData(pd.DataFrame(records), np.asarray(responses), np.asarray(contexts),
                        np.asarray(tv), np.asarray(cv), metadata).validate()
    data.save(output)
    return data