File size: 13,682 Bytes
b04653d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c3b49d6
b04653d
 
c3b49d6
b04653d
c3b49d6
b04653d
 
 
 
9c88c1e
b04653d
 
 
 
 
 
 
 
 
c3b49d6
b04653d
 
 
 
 
 
 
 
 
 
 
 
 
c3b49d6
b04653d
 
 
 
 
 
 
c3b49d6
b04653d
 
 
 
 
 
 
 
c3b49d6
b04653d
 
 
c3b49d6
b04653d
 
c3b49d6
b04653d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c3b49d6
25b085e
b04653d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bbba11d
 
 
 
 
c3b49d6
bbba11d
b04653d
 
 
 
 
c3b49d6
b04653d
 
c3b49d6
b04653d
 
 
 
 
 
 
 
 
 
 
 
 
 
c3b49d6
b04653d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
526f1d2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c3b49d6
526f1d2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c3b49d6
526f1d2
 
c3b49d6
526f1d2
c3b49d6
526f1d2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
"""
Cross-dataset early integration (ADR-0001, Mode A — T7).

Build ONE combined AnnData from >=2 registered datasets by:
  1. harmonising each to a shared gene-symbol feature axis,
  2. intersecting features (join="inner"), and
  3. concatenating samples with a ``batch`` obs column = dataset_id.

The result is a *standard* AnnData that the existing single-dataset tools
(``decoupler_differential_expression`` and the enrichment tools) consume
unchanged — Mode A pools the matrices and a downstream DE models ``dataset_id``
as a batch covariate (see ``decoupler_differential_expression``'s
``batch_column``, T8). Whether a set of datasets is eligible for early pooling is
decided upstream by ``get_integration_plan`` (mode == "early"); this module
performs the mechanics once that verdict is in (the gating lives in the
``decoupler_integrate_datasets`` tool, T9).

Pure core: ``combine_anndatas(adatas, ...)`` operates on already-loaded AnnData
objects (unit-testable, no network). ``build_combined_anndata(dataset_ids, ...)``
is the thin loader wrapper that resolves each manifest's hosted h5ad.

v1 scope / limitations:
- Harmonises to gene symbols: a dataset whose ``feature_id_type`` is
  "gene_symbol" uses var_names directly; otherwise a var column named
  SYMBOL/gene_symbol/gene_name is used if present; anything else raises (probe
  collapse / ortholog mapping / entrez without a symbol column are NOT handled
  here — such datasets do not reach an ``early`` verdict in v1).
- Duplicate symbols within a dataset are collapsed keep-first (preserves raw
  integer counts for the DESeq2 path — no averaging).
"""

from __future__ import annotations

from collections.abc import Sequence
from pathlib import Path
from typing import TYPE_CHECKING

import anndata as ad
import numpy as np

from src.core.paths import OUTPUT_DIR

if TYPE_CHECKING:
    from anndata import AnnData

# var columns checked, in order, for a gene-symbol axis when feature_id_type
# is not already "gene_symbol".
_SYMBOL_VAR_COLUMNS = ("SYMBOL", "gene_symbol", "Gene_Symbol", "symbol", "gene_name")


def _symbol_axis(adata: AnnData, feature_id_type: str) -> list[str]:
    """Return a per-var list of gene symbols for ``adata``, or raise if unavailable."""
    if feature_id_type == "gene_symbol":
        return [str(v) for v in adata.var_names]
    for col in _SYMBOL_VAR_COLUMNS:
        if col in adata.var.columns:
            return [str(v) for v in adata.var[col].tolist()]
    raise ValueError(
        f"cannot resolve gene symbols: feature_id_type='{feature_id_type}' and no "
        f"symbol column {list(_SYMBOL_VAR_COLUMNS)} in var. Mode-A v1 needs a "
        f"gene-symbol axis (probe collapse / id-mapping is not handled here)."
    )


def _to_symbol_adata(adata: AnnData, feature_id_type: str) -> AnnData:
    """Return a copy of ``adata`` whose var_names are unique gene symbols.

    Drops features with empty/missing symbols; collapses duplicate symbols
    keep-first (preserves raw counts — no averaging)."""
    symbols = _symbol_axis(adata, feature_id_type)
    a = adata.copy()
    a.var_names = [str(s) for s in symbols]
    keep = np.array([bool(s) and s.strip().lower() not in ("nan", "none", "") for s in a.var_names])
    a = a[:, keep].copy()
    # keep-first on duplicate symbols (np.unique returns the first index per value)
    _, first_idx = np.unique(np.asarray(a.var_names, dtype=object), return_index=True)
    a = a[:, np.sort(first_idx)].copy()
    return a


def combine_anndatas(
    adatas: Sequence[AnnData],
    dataset_ids: Sequence[str],
    feature_id_types: Sequence[str],
    *,
    design_factor: str | None = None,
    batch_key: str = "batch",
    min_shared_features: int = 200,
) -> AnnData:
    """Pure core: feature-intersection concat of >=2 AnnData into one combined AnnData.

    Each input is harmonised to a unique gene-symbol var axis, the var
    intersection is taken (``join="inner"``), a ``batch_key`` obs column is set to
    the dataset_id, and samples are concatenated (obs names disambiguated per
    batch with a ``-<dataset_id>`` suffix).

    Raises
    ------
    ValueError
        on <2 inputs, mismatched argument lengths, an unresolvable symbol axis,
        a missing ``design_factor`` column in any dataset, or a feature
        intersection smaller than ``min_shared_features``.
    """
    if len(adatas) < 2:
        raise ValueError(f"early integration needs >=2 datasets; got {len(adatas)}")
    if not (len(adatas) == len(dataset_ids) == len(feature_id_types)):
        raise ValueError(
            "adatas, dataset_ids, and feature_id_types must be the same length "
            f"({len(adatas)}, {len(dataset_ids)}, {len(feature_id_types)})"
        )

    prepared: list[AnnData] = []
    for a, did, fit in zip(adatas, dataset_ids, feature_id_types, strict=True):
        sa = _to_symbol_adata(a, fit)
        if design_factor is not None and design_factor not in sa.obs.columns:
            raise ValueError(
                f"dataset '{did}' has no obs column '{design_factor}' required for "
                f"the requested contrast."
            )
        prepared.append(sa)

    shared = set(prepared[0].var_names)
    for sa in prepared[1:]:
        shared &= set(sa.var_names)
    if len(shared) < min_shared_features:
        raise ValueError(
            f"only {len(shared)} shared gene symbols across {list(dataset_ids)} "
            f"(minimum {min_shared_features}); the datasets likely do not share a "
            f"feature space and are not eligible for early pooling."
        )

    combined = ad.concat(
        prepared,
        axis=0,
        join="inner",  # feature intersection
        label=batch_key,
        keys=[str(d) for d in dataset_ids],
        index_unique="-",  # disambiguate duplicate sample barcodes across cohorts
        merge="same",
    )
    combined.obs[batch_key] = combined.obs[batch_key].astype("category")
    return combined


def _resolve_to_local(url_or_path: str) -> tuple[str, bool]:
    """Return (local_path, is_temp). Downloads http(s)/ftp URLs to a temp file.

    Delegates to the shared authenticated resolver so private huggingface.co
    files download with HF_TOKEN (see src/core/data_io.py)."""
    from src.core.data_io import resolve_to_local_path

    return resolve_to_local_path(url_or_path)


def build_combined_anndata(
    dataset_ids: Sequence[str],
    *,
    design_factor: str | None = None,
    batch_key: str = "batch",
    min_shared_features: int = 200,
    out_prefix: str | None = None,
) -> dict:
    """Loader wrapper: resolve each registered dataset's hosted h5ad and combine.

    Resolves each ``dataset_id`` to its manifest's ``expression_source.url``,
    reads the AnnData, and delegates to :func:`combine_anndatas`. Writes the
    combined AnnData to ``OUTPUT_DIR`` and returns a summary dict.
    """
    import scanpy as sc

    from src.datasets.registry import load_manifest

    if len(dataset_ids) < 2:
        raise ValueError(f"early integration needs >=2 datasets; got {len(dataset_ids)}")

    adatas: list[AnnData] = []
    feature_id_types: list[str] = []
    temps: list[str] = []
    try:
        for did in dataset_ids:
            m = load_manifest(did)
            url = (m.expression_source or {}).get("url")
            if not url:
                raise ValueError(f"dataset '{did}' has no expression_source.url")
            local, is_temp = _resolve_to_local(url)
            if is_temp:
                temps.append(local)
            adatas.append(sc.read_h5ad(local))
            feature_id_types.append(m.feature_id_type)

        combined = combine_anndatas(
            adatas,
            dataset_ids,
            feature_id_types,
            design_factor=design_factor,
            batch_key=batch_key,
            min_shared_features=min_shared_features,
        )
    finally:
        for t in temps:
            Path(t).unlink(missing_ok=True)

    out_prefix = out_prefix or ("combined_" + "_".join(str(d) for d in dataset_ids))[:80]
    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
    out_path = OUTPUT_DIR / f"{out_prefix}.h5ad"
    combined.write_h5ad(out_path)

    per_batch = combined.obs[batch_key].value_counts().to_dict()
    return {
        "output_path": str(out_path.resolve()),
        "n_obs": int(combined.n_obs),
        "n_vars": int(combined.n_vars),
        "batch_key": batch_key,
        "shared_feature_space": "gene_symbol",
        "per_batch_n": {str(k): int(v) for k, v in per_batch.items()},
        "dataset_ids": list(dataset_ids),
    }


# ---------------------------------------------------------------------------
# ComBat pre-correction for the per-sample activity-scoring path (ADR-0001 #10)
# ---------------------------------------------------------------------------
#
# The DE path models `dataset` as a covariate (decoupler_integrate_datasets). The
# per-sample activity-scoring path has no design matrix to carry one, so the pooled
# matrix is batch-corrected with ComBat BEFORE scoring. Standard ComBat (empirical
# Bayes, scanpy.pp.combat) on the log-normalised matrix is used — NOT ComBat-seq —
# because the scoring path log-normalises anyway and the corrected matrix never feeds
# a count model (that path uses the covariate route). See ADR-0001 item 10.
#
# Data levels that ARE log-scale already (ComBat applied directly, no transform).
_LOG_SCALE_LEVELS = ("log_expression", "log_ratio")
# Linear-but-not-count levels: log1p before ComBat.
_LINEAR_LEVELS = ("tpm", "fpkm")


def _looks_like_raw_counts(X) -> bool:
    """Heuristic: does this matrix look like raw integer counts?

    Non-negative, (near-)integer, and a large dynamic range — the same signal the
    DE tool uses to guard DESeq2. Used as a safety net when the declared
    ``data_level`` is missing or disagrees with the values."""
    import numpy as np

    arr = np.asarray(X.toarray() if hasattr(X, "toarray") else X, dtype=float)
    sample = arr[: min(50, arr.shape[0])]
    finite = sample[np.isfinite(sample)]
    if finite.size == 0:
        return False
    return (
        float(finite.min()) >= 0
        and float(finite.max()) > 30
        and bool(np.allclose(finite, np.round(finite), atol=1e-6))
    )


def _genes_variable_within_each_batch(adata: AnnData, batch_key: str, tol: float = 1e-12):
    """Boolean mask of genes with non-zero variance within EVERY batch.

    ComBat estimates a per-batch location/scale per gene; a gene that is constant
    inside any batch yields a zero-variance estimate and NaNs. Drop those first."""
    import numpy as np

    X = np.asarray(adata.X.toarray() if hasattr(adata.X, "toarray") else adata.X, dtype=float)
    keep = np.ones(adata.n_vars, dtype=bool)
    batches = adata.obs[batch_key].astype(str).to_numpy()
    for b in np.unique(batches):
        sub = X[batches == b]
        keep &= sub.std(axis=0) > tol
    return keep


def batch_correct_for_scoring(
    adata: AnnData,
    *,
    batch_key: str = "batch",
    data_level: str | None = None,
    min_per_batch: int = 2,
) -> tuple[AnnData, dict]:
    """ComBat-correct a pooled AnnData for per-sample activity scoring.

    Brings the matrix to a log scale appropriate for decoupleR scoring (keyed on
    ``data_level``, with a raw-counts safety detector), drops genes that are
    constant within any batch, then applies ``scanpy.pp.combat`` keyed on
    ``batch_key``. Returns ``(corrected_adata, info)``.

    This is the **no-design-matrix** counterpart to the DE covariate route — use it
    only for per-sample scoring / clustering / PCA, NEVER as a pre-step to DE
    testing (ComBat + naive DE inflates false positives; ADR-0001). ComBat runs
    without a biological covariate, so dataset-level structure is removed wholesale.

    Raises
    ------
    ValueError
        when ``batch_key`` is absent, a batch has < ``min_per_batch`` samples
        (ComBat cannot estimate it), or < 2 non-constant genes remain.
    """
    import numpy as np
    import scanpy as sc

    if batch_key not in adata.obs.columns:
        raise ValueError(f"batch_key '{batch_key}' not found in obs.")

    counts = adata.obs[batch_key].value_counts()
    too_small = counts[counts < min_per_batch]
    if len(too_small):
        raise ValueError(
            f"ComBat needs >= {min_per_batch} samples per batch; too small: "
            f"{ {str(k): int(v) for k, v in too_small.items()} }."
        )

    a = adata.copy()
    a.X = np.asarray(a.X.toarray() if hasattr(a.X, "toarray") else a.X, dtype=np.float64)

    level = (data_level or "").lower()
    if level == "raw_counts" or (level not in _LOG_SCALE_LEVELS and _looks_like_raw_counts(a.X)):
        sc.pp.normalize_total(a, target_sum=1e4)
        sc.pp.log1p(a)
        normalization = "normalize_total(1e4)+log1p"
    elif level in _LINEAR_LEVELS:
        sc.pp.log1p(a)
        normalization = "log1p"
    else:
        normalization = "none (already log-scale)"

    keep = _genes_variable_within_each_batch(a, batch_key)
    n_dropped = int((~keep).sum())
    if n_dropped:
        a = a[:, keep].copy()
    if a.n_vars < 2:
        raise ValueError(
            "fewer than 2 genes are non-constant within every batch; cannot ComBat-"
            "correct (the cohorts likely share almost no usable feature space)."
        )

    sc.pp.combat(a, key=batch_key)

    info = {
        "method": "ComBat (scanpy.pp.combat)",
        "normalization": normalization,
        "n_genes_corrected": int(a.n_vars),
        "n_genes_dropped_constant_within_batch": n_dropped,
        "batch_key": batch_key,
        "data_level": data_level,
    }
    return a, info