vima-spatial 0.2.2__tar.gz → 0.2.3__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/PKG-INFO +5 -2
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/README.md +4 -1
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/pyproject.toml +4 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/setup.cfg +1 -1
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/__init__.py +2 -2
- vima_spatial-0.2.3/src/vima/cc.py +465 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/data/patchcollection.py +101 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/data/samples.py +45 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/fingerprints.py +88 -1
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/ingest/dimreduce.py +77 -1
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/ingest/ingest.py +103 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/ingest/nonst.py +71 -2
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/ingest/st.py +150 -14
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/ingest/util.py +16 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/models/__init__.py +23 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/models/simple_vae.py +3 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/models/vae.py +19 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/patchfeatures.py +105 -58
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/train/logging.py +3 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/train/training.py +112 -1
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/vis/features.py +46 -32
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/vis/patches.py +129 -59
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/vis/spatial.py +69 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/vis/umaps.py +15 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima_spatial.egg-info/PKG-INFO +5 -2
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima_spatial.egg-info/SOURCES.txt +2 -2
- vima_spatial-0.2.3/tests/test_ra_regression.py +69 -0
- vima_spatial-0.2.2/MANIFEST.in +0 -1
- vima_spatial-0.2.2/src/vima/cc.py +0 -199
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/data/__init__.py +0 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/ingest/__init__.py +0 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/models/resnet_vae.py +0 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/models/resnetlight_decoder.py +0 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/models/resnetlight_encoder.py +0 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/train/__init__.py +0 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/vis/__init__.py +0 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/vis/patchexamples.py +0 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima_spatial.egg-info/dependency_links.txt +0 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima_spatial.egg-info/requires.txt +0 -0
- {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima_spatial.egg-info/top_level.txt +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: vima-spatial
|
|
3
|
-
Version: 0.2.
|
|
3
|
+
Version: 0.2.3
|
|
4
4
|
Summary: variational inference-based microniche analysis
|
|
5
5
|
Home-page: https://github.com/yakirr/vima
|
|
6
6
|
Author: Yakir Reshef
|
|
@@ -39,9 +39,12 @@ Variational inference-based microniche analysis is a method for conducting case-
|
|
|
39
39
|
```
|
|
40
40
|
pip install vima-spatial
|
|
41
41
|
```
|
|
42
|
+
Note that `vima` requires `pytorch` and `harmonypy`. These should install automatically through `pip`, but if you have trouble, try installing them first, verifying that they work, and then installing `vima`.
|
|
42
43
|
|
|
43
44
|
## demo
|
|
44
|
-
|
|
45
|
+
To get started with an example analysis on a toy spatial transcriptomics dataset, take a look at our brief demo. You can see a [completed read-only version](https://github.com/yakirr/vima/blob/main/demo/demo_ST_minimal.ipynb) or run an [interactive version](https://colab.research.google.com/github/yakirr/tpae/blob/main/demo/demo_ST_minimal.ipynb) yourself on a Google Colab GPU (though this requires Colab Pro due to insufficient memory provided on the free tier.)
|
|
46
|
+
|
|
47
|
+
To see how to apply `vima` to a stain-based modality like CODEX, immunohistochemistry, or immunofluorescence, look at our [immunofluorescence demo](https://github.com/yakirr/vima/blob/main/demo/demo_IF.ipynb).
|
|
45
48
|
|
|
46
49
|
## citation
|
|
47
50
|
If you use `vima`, please cite:
|
|
@@ -6,9 +6,12 @@ Variational inference-based microniche analysis is a method for conducting case-
|
|
|
6
6
|
```
|
|
7
7
|
pip install vima-spatial
|
|
8
8
|
```
|
|
9
|
+
Note that `vima` requires `pytorch` and `harmonypy`. These should install automatically through `pip`, but if you have trouble, try installing them first, verifying that they work, and then installing `vima`.
|
|
9
10
|
|
|
10
11
|
## demo
|
|
11
|
-
|
|
12
|
+
To get started with an example analysis on a toy spatial transcriptomics dataset, take a look at our brief demo. You can see a [completed read-only version](https://github.com/yakirr/vima/blob/main/demo/demo_ST_minimal.ipynb) or run an [interactive version](https://colab.research.google.com/github/yakirr/tpae/blob/main/demo/demo_ST_minimal.ipynb) yourself on a Google Colab GPU (though this requires Colab Pro due to insufficient memory provided on the free tier.)
|
|
13
|
+
|
|
14
|
+
To see how to apply `vima` to a stain-based modality like CODEX, immunohistochemistry, or immunofluorescence, look at our [immunofluorescence demo](https://github.com/yakirr/vima/blob/main/demo/demo_IF.ipynb).
|
|
12
15
|
|
|
13
16
|
## citation
|
|
14
17
|
If you use `vima`, please cite:
|
|
@@ -8,11 +8,11 @@ from . import models
|
|
|
8
8
|
from .data.patchcollection import PatchCollection
|
|
9
9
|
from .data.samples import read_samples, reindex_by_sid
|
|
10
10
|
from .train.training import train, fit, set_seed
|
|
11
|
-
from .cc import latentreps, association
|
|
11
|
+
from .cc import latentreps, association, compute_mams
|
|
12
12
|
from .fingerprints import Fingerprints
|
|
13
13
|
from .patchfeatures import cell_type_counts, expression_profiles, test_features
|
|
14
14
|
|
|
15
15
|
__all__ = ['d', 'cc', 't', 'pp', 'v', 'models',
|
|
16
16
|
'PatchCollection', 'read_samples', 'reindex_by_sid',
|
|
17
|
-
'train', 'fit', 'set_seed', 'latentreps', 'association',
|
|
17
|
+
'train', 'fit', 'set_seed', 'latentreps', 'association', 'compute_mams',
|
|
18
18
|
'Fingerprints', 'cc', 'cell_type_counts', 'expression_profiles', 'test_features']
|
|
@@ -0,0 +1,465 @@
|
|
|
1
|
+
import os
|
|
2
|
+
from concurrent.futures import ThreadPoolExecutor
|
|
3
|
+
import numpy as np
|
|
4
|
+
import scanpy as sc
|
|
5
|
+
import anndata as ad
|
|
6
|
+
import pandas as pd
|
|
7
|
+
import torch
|
|
8
|
+
from torch.utils.data import DataLoader
|
|
9
|
+
import cna
|
|
10
|
+
import warnings
|
|
11
|
+
import scipy.sparse as sp
|
|
12
|
+
import scipy.stats as st
|
|
13
|
+
from argparse import Namespace
|
|
14
|
+
from tqdm import tqdm
|
|
15
|
+
from .fingerprints import Fingerprints
|
|
16
|
+
pb = lambda x: tqdm(x, ncols=100)
|
|
17
|
+
|
|
18
|
+
def anndata(patchmeta, Z, var_names=None, use_rep='X', n_comps=10, **kwargs):
|
|
19
|
+
"""
|
|
20
|
+
Build an AnnData with a neighbor graph from an embedding matrix.
|
|
21
|
+
|
|
22
|
+
Parameters
|
|
23
|
+
----------
|
|
24
|
+
patchmeta
|
|
25
|
+
Per-patch metadata to store in ``.obs``.
|
|
26
|
+
Z
|
|
27
|
+
Patch-by-dimension embedding matrix.
|
|
28
|
+
use_rep
|
|
29
|
+
Representation used to build the neighbor graph: 'X' (the embedding) or
|
|
30
|
+
'X_pca'.
|
|
31
|
+
n_comps
|
|
32
|
+
Number of PCs when ``use_rep='X_pca'``.
|
|
33
|
+
|
|
34
|
+
Returns
|
|
35
|
+
-------
|
|
36
|
+
AnnData
|
|
37
|
+
Embedding in ``X``, patch metadata in ``.obs``, and a nearest-neighbor
|
|
38
|
+
graph.
|
|
39
|
+
"""
|
|
40
|
+
d = ad.AnnData(Z)
|
|
41
|
+
if var_names is not None:
|
|
42
|
+
d.var_names = var_names
|
|
43
|
+
obs = patchmeta.copy()
|
|
44
|
+
obs.index = obs.index.astype(str)
|
|
45
|
+
d.obs = obs
|
|
46
|
+
|
|
47
|
+
if use_rep == 'X_pca':
|
|
48
|
+
sc.tl.pca(d, n_comps=min(n_comps, Z.shape[1]-1))
|
|
49
|
+
|
|
50
|
+
sc.pp.neighbors(d, use_rep=use_rep, **kwargs)
|
|
51
|
+
|
|
52
|
+
return d
|
|
53
|
+
|
|
54
|
+
def apply(models, P, batch_size=1000, with_mse=False):
|
|
55
|
+
"""
|
|
56
|
+
Run a trained model ensemble over a PatchCollection to embed every patch.
|
|
57
|
+
|
|
58
|
+
Runs with augmentation disabled and each model in eval mode.
|
|
59
|
+
|
|
60
|
+
Parameters
|
|
61
|
+
----------
|
|
62
|
+
models
|
|
63
|
+
Trained ensemble from `cVAE`.
|
|
64
|
+
P
|
|
65
|
+
Patches to embed.
|
|
66
|
+
with_mse
|
|
67
|
+
Also return per-patch, per-channel reconstruction error.
|
|
68
|
+
|
|
69
|
+
Returns
|
|
70
|
+
-------
|
|
71
|
+
ndarray or tuple
|
|
72
|
+
Model-by-patch-by-dimension embeddings; if `with_mse`, also a
|
|
73
|
+
model-by-patch-by-marker array of reconstruction MSEs.
|
|
74
|
+
"""
|
|
75
|
+
P.pytorch_mode()
|
|
76
|
+
P.augmentation_off()
|
|
77
|
+
for model in models:
|
|
78
|
+
model.eval()
|
|
79
|
+
|
|
80
|
+
eval_loader = DataLoader(
|
|
81
|
+
dataset=P,
|
|
82
|
+
batch_size=batch_size,
|
|
83
|
+
shuffle=False)
|
|
84
|
+
|
|
85
|
+
Zs = {modelid: [] for modelid in range(len(models))}
|
|
86
|
+
if with_mse:
|
|
87
|
+
MSEs = {modelid: [] for modelid in range(len(models))}
|
|
88
|
+
with torch.no_grad():
|
|
89
|
+
for batch in pb(eval_loader):
|
|
90
|
+
for modelid, model in enumerate(models):
|
|
91
|
+
if with_mse:
|
|
92
|
+
x_recon, mean, _ = model(batch, sample_from_latent=False)
|
|
93
|
+
Zs[modelid].append(mean.reshape(len(batch[0]), -1).detach().cpu().numpy())
|
|
94
|
+
MSEs[modelid].append(((x_recon - batch[0]) ** 2).mean(dim=(2, 3)).detach().cpu().numpy())
|
|
95
|
+
else:
|
|
96
|
+
Zs[modelid].append(model.embedding(batch).detach().cpu().numpy())
|
|
97
|
+
|
|
98
|
+
Zs_out = np.array([np.concatenate(Z) for Z in Zs.values()])
|
|
99
|
+
if with_mse:
|
|
100
|
+
MSEs_out = np.array([np.concatenate(M) for M in MSEs.values()])
|
|
101
|
+
return Zs_out, MSEs_out
|
|
102
|
+
return Zs_out
|
|
103
|
+
|
|
104
|
+
def latentreps(models, P, use_rep='X', n_comps=100, with_mse=True, **kwargs):
|
|
105
|
+
"""
|
|
106
|
+
Compute patch fingerprints from a trained model ensemble.
|
|
107
|
+
|
|
108
|
+
Applies each model to every patch to obtain its latent embedding, then
|
|
109
|
+
builds a per-model nearest-neighbor graph over the patches. These are the
|
|
110
|
+
"fingerprints" consumed by `association`.
|
|
111
|
+
|
|
112
|
+
Parameters
|
|
113
|
+
----------
|
|
114
|
+
models
|
|
115
|
+
Trained ensemble from `cVAE`.
|
|
116
|
+
P
|
|
117
|
+
Patches to embed (typically the refined collection).
|
|
118
|
+
use_rep
|
|
119
|
+
Representation used to build the neighbor graph: 'X' (the embedding) or
|
|
120
|
+
'X_pca'.
|
|
121
|
+
n_comps
|
|
122
|
+
Number of PCs when ``use_rep='X_pca'``.
|
|
123
|
+
with_mse
|
|
124
|
+
Also store per-patch and per-channel reconstruction error.
|
|
125
|
+
|
|
126
|
+
Returns
|
|
127
|
+
-------
|
|
128
|
+
Fingerprints
|
|
129
|
+
Per-model embeddings and neighbor graphs, with ``.obs`` carrying patch
|
|
130
|
+
metadata.
|
|
131
|
+
"""
|
|
132
|
+
print('applying models')
|
|
133
|
+
result = apply(models, P, with_mse=with_mse)
|
|
134
|
+
Zs, MSEs = result if with_mse else (result, None)
|
|
135
|
+
|
|
136
|
+
print('computing nearest-neighbor graphs')
|
|
137
|
+
ds = [anndata(P.meta, Z, use_rep=use_rep, n_comps=n_comps, **kwargs) for Z in pb(Zs)]
|
|
138
|
+
fp = Fingerprints.from_list(ds)
|
|
139
|
+
|
|
140
|
+
if MSEs is not None:
|
|
141
|
+
mean_mse = MSEs.mean(axis=0)
|
|
142
|
+
marker_names = next(iter(P.samples.values())).coords['marker'].values
|
|
143
|
+
per_channel = pd.DataFrame(mean_mse, index=fp.obs.index, columns=marker_names)
|
|
144
|
+
fp.obsm['per_channel_mse'] = per_channel
|
|
145
|
+
fp.obs['mse'] = per_channel.mean(axis=1)
|
|
146
|
+
return fp
|
|
147
|
+
|
|
148
|
+
|
|
149
|
+
def _tail_counts_total(znull, t2_sorted, nthreads):
|
|
150
|
+
"""Total count of znull**2 entries >= each (sorted) squared threshold.
|
|
151
|
+
"""
|
|
152
|
+
col_chunks = np.array_split(np.arange(znull.shape[1]), nthreads)
|
|
153
|
+
|
|
154
|
+
def chunk_hist(cols):
|
|
155
|
+
sub = np.square(znull[:, cols]).ravel()
|
|
156
|
+
# number of thresholds each entry exceeds, in [0, len(t2_sorted)]
|
|
157
|
+
m = np.searchsorted(t2_sorted, sub, side="right")
|
|
158
|
+
return np.bincount(m, minlength=t2_sorted.size + 1)
|
|
159
|
+
|
|
160
|
+
with ThreadPoolExecutor(nthreads) as ex:
|
|
161
|
+
hist = np.sum(list(ex.map(chunk_hist, col_chunks)), axis=0)
|
|
162
|
+
|
|
163
|
+
# counts_sorted[i] = # entries exceeding at least (i+1) thresholds
|
|
164
|
+
return np.cumsum(hist[::-1])[::-1][1:]
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
def empirical_fdrs(z, znull, thresholds):
|
|
168
|
+
"""
|
|
169
|
+
Compute the empirical FDR at each threshold from observed and null statistics.
|
|
170
|
+
|
|
171
|
+
At each threshold the FDR estimate is the mean number of null statistics
|
|
172
|
+
exceeding it (per null realization) divided by the number of observed
|
|
173
|
+
statistics exceeding it, comparing magnitudes.
|
|
174
|
+
|
|
175
|
+
Parameters
|
|
176
|
+
----------
|
|
177
|
+
z
|
|
178
|
+
Observed per-microniche statistics.
|
|
179
|
+
znull
|
|
180
|
+
Null statistics, one column per permutation (aligned with `z`).
|
|
181
|
+
thresholds
|
|
182
|
+
Thresholds at which to evaluate the FDR.
|
|
183
|
+
|
|
184
|
+
Returns
|
|
185
|
+
-------
|
|
186
|
+
ndarray
|
|
187
|
+
Estimated FDR at each threshold.
|
|
188
|
+
"""
|
|
189
|
+
if znull.shape[0] != len(z):
|
|
190
|
+
raise ValueError("shape mismatch")
|
|
191
|
+
|
|
192
|
+
if znull.ndim == 1:
|
|
193
|
+
znull = znull[:, None]
|
|
194
|
+
ncols = znull.shape[1]
|
|
195
|
+
|
|
196
|
+
t2 = np.square(np.asarray(thresholds, dtype=float))
|
|
197
|
+
order = np.argsort(t2)
|
|
198
|
+
|
|
199
|
+
nthreads = min(ncols, (os.cpu_count() or 1))
|
|
200
|
+
counts_sorted = _tail_counts_total(znull, t2[order], nthreads)
|
|
201
|
+
mean_tails = np.empty(t2.shape)
|
|
202
|
+
mean_tails[order] = counts_sorted / ncols
|
|
203
|
+
|
|
204
|
+
ranks = len(z) - np.searchsorted(np.sort(np.square(z)), t2, side="left")
|
|
205
|
+
|
|
206
|
+
return mean_tails / ranks
|
|
207
|
+
|
|
208
|
+
|
|
209
|
+
def _power_ratio(x, power, axis):
|
|
210
|
+
"""Meta-analysis power ratio (x**power).sum / (x**2).sum along `axis` (the model axis)."""
|
|
211
|
+
return (x**power).sum(axis=axis) / (x**2).sum(axis=axis)
|
|
212
|
+
|
|
213
|
+
|
|
214
|
+
def _association(MAMresid, M, y, batches, donorids, rng, Nnull=10_000,
|
|
215
|
+
max_num_mns=5_000, show_progress=False):
|
|
216
|
+
"""
|
|
217
|
+
Run the permutation-based microniche association test.
|
|
218
|
+
|
|
219
|
+
Correlates each microniche's residualized cross-sample abundance with the
|
|
220
|
+
covariate-conditioned phenotype, meta-analyzing across models, and calibrates
|
|
221
|
+
both a global p-value and per-microniche FDRs against permutation nulls. Null
|
|
222
|
+
phenotypes are drawn at the donor level when `donorids` is given, or within
|
|
223
|
+
`batches` if that is given, or unconditionally otherwise.
|
|
224
|
+
|
|
225
|
+
Parameters
|
|
226
|
+
----------
|
|
227
|
+
MAMresid
|
|
228
|
+
Residualized abundances, shaped sample-by-model-by-microniche.
|
|
229
|
+
M
|
|
230
|
+
Residualization matrix (sample-by-sample) used to residualize any covariates
|
|
231
|
+
out of MAMresid.
|
|
232
|
+
y
|
|
233
|
+
Sample-level phenotype.
|
|
234
|
+
max_num_mns
|
|
235
|
+
Cap on microniches subsampled for the global and FDR computations.
|
|
236
|
+
|
|
237
|
+
Returns
|
|
238
|
+
-------
|
|
239
|
+
Namespace
|
|
240
|
+
Global p-value, per-microniche and per-model coefficients, mixing
|
|
241
|
+
weights, FDR table, and the associated null distributions.
|
|
242
|
+
"""
|
|
243
|
+
# prep data
|
|
244
|
+
y = (y - y.mean())/y.std()
|
|
245
|
+
n = len(y)
|
|
246
|
+
ycond = M.dot(y)
|
|
247
|
+
ycond /= ycond.std(axis=0)
|
|
248
|
+
|
|
249
|
+
# make null phenotypes
|
|
250
|
+
if donorids is not None:
|
|
251
|
+
y_ = cna.tl._stats.grouplevel_permutation(donorids, y, Nnull)
|
|
252
|
+
else:
|
|
253
|
+
y_ = cna.tl._stats.conditional_permutation(batches, y, Nnull)
|
|
254
|
+
ycond_ = M.dot(y_)
|
|
255
|
+
ycond_ /= ycond_.std(axis=0)
|
|
256
|
+
|
|
257
|
+
# get microniche coefficients and weights (over all patches)
|
|
258
|
+
mncorrs = (ycond[:,None,None]*MAMresid).mean(axis=0)
|
|
259
|
+
weights = (mncorrs**2) / (mncorrs**2).sum(axis=0)
|
|
260
|
+
mncorrs_meta = _power_ratio(mncorrs, 3, axis=0)
|
|
261
|
+
|
|
262
|
+
# subsample patches (last axis of MAMresid) for the expensive global/FDR machinery
|
|
263
|
+
Npatches = MAMresid.shape[2]
|
|
264
|
+
if Npatches > max_num_mns:
|
|
265
|
+
sub = rng.choice(Npatches, size=max_num_mns, replace=False)
|
|
266
|
+
else:
|
|
267
|
+
sub = np.arange(Npatches)
|
|
268
|
+
MAMresid_sub = MAMresid[:, :, sub]
|
|
269
|
+
mncorrs_sub = mncorrs[:, sub]
|
|
270
|
+
mncorrs_meta_sub = mncorrs_meta[sub]
|
|
271
|
+
|
|
272
|
+
# meta-analyzed mn coefficients and global test statistics (on the subsampled patches).
|
|
273
|
+
# We loop over the (few) models and accumulate the per-model power sums Sq = sum_m nm_m**q
|
|
274
|
+
# incrementally, so we never materialize the full (Nnull, Nmodels, max_num_mns) array.
|
|
275
|
+
globalstat = _power_ratio(mncorrs_sub, 4, axis=0).mean()
|
|
276
|
+
ycond_T = np.ascontiguousarray(ycond_.T, dtype=np.float32) # (Nnull, n)
|
|
277
|
+
MAMresid_sub32 = MAMresid_sub.astype(np.float32)
|
|
278
|
+
S2 = np.zeros((ycond_.shape[1], MAMresid_sub.shape[2]), dtype=np.float32) # (Nnull, max_num_mns)
|
|
279
|
+
S3 = np.zeros_like(S2); S4 = np.zeros_like(S2)
|
|
280
|
+
MAMresid_sub32_permodel = [np.ascontiguousarray(MAMresid_sub32[:, m, :]) for m in range(MAMresid_sub32.shape[1])]
|
|
281
|
+
for myMAM in MAMresid_sub32_permodel:
|
|
282
|
+
nm = (ycond_T @ myMAM) / n # null mn coeffs for model m: (Nnull, npatch)
|
|
283
|
+
nm2 = nm * nm
|
|
284
|
+
S2 += nm2; S3 += nm2 * nm; S4 += nm2 * nm2
|
|
285
|
+
|
|
286
|
+
nullglobalstats = (S4 / S2).mean(axis=1)
|
|
287
|
+
nullmncorrs_meta = (S3 / S2).T
|
|
288
|
+
|
|
289
|
+
# compute global p-vaule
|
|
290
|
+
p = ((nullglobalstats >= globalstat).sum() + 1)/(len(nullglobalstats) + 1)
|
|
291
|
+
print(f'\033[32mP = {p}\033[0m')
|
|
292
|
+
if p <= 1/(Nnull + 1)+1e-10:
|
|
293
|
+
warnings.warn('global association p-value attained minimal possible value. '+\
|
|
294
|
+
'Consider increasing Nnull')
|
|
295
|
+
|
|
296
|
+
thr = np.quantile(np.abs(mncorrs_meta_sub), np.arange(0.01, 1, 0.01))
|
|
297
|
+
fdrs = empirical_fdrs(mncorrs_meta_sub, nullmncorrs_meta, thr)
|
|
298
|
+
fdrs = pd.DataFrame({
|
|
299
|
+
'threshold':thr,
|
|
300
|
+
'fdr':fdrs})
|
|
301
|
+
|
|
302
|
+
res = {'p':p, 'mncorrs':mncorrs_meta, 'fdrs':fdrs,
|
|
303
|
+
'globalstat':globalstat, 'nullglobalstats':nullglobalstats,
|
|
304
|
+
'weights':weights,
|
|
305
|
+
'nullmncorrs':nullmncorrs_meta,
|
|
306
|
+
'permodel_mncorrs':mncorrs,
|
|
307
|
+
'MAMres':MAMresid,
|
|
308
|
+
'ycond':ycond
|
|
309
|
+
}
|
|
310
|
+
|
|
311
|
+
return Namespace(**res)
|
|
312
|
+
|
|
313
|
+
|
|
314
|
+
def compute_mams(ds, sid_name, nsteps=None, self_weight=1, show_progress=False):
|
|
315
|
+
"""
|
|
316
|
+
Precompute per-model sample-by-microniche abundance matrices.
|
|
317
|
+
|
|
318
|
+
Runs only the expensive graph-diffusion step, independent of any phenotype,
|
|
319
|
+
batch, or covariate choice. Pass the result to ``association(..., MAMs=...)``
|
|
320
|
+
to reuse it across multiple phenotypes without recomputing.
|
|
321
|
+
|
|
322
|
+
Parameters
|
|
323
|
+
----------
|
|
324
|
+
ds
|
|
325
|
+
Fingerprints from `latentreps`.
|
|
326
|
+
sid_name
|
|
327
|
+
Column in ``ds.obs`` giving each patch's sample ID.
|
|
328
|
+
|
|
329
|
+
Returns
|
|
330
|
+
-------
|
|
331
|
+
list
|
|
332
|
+
One abundance matrix per model.
|
|
333
|
+
"""
|
|
334
|
+
print('computing MAT') #TODO: rename MAM to MAT in code if we keep this nomenclature
|
|
335
|
+
MAMs = []
|
|
336
|
+
for d in tqdm(ds.modelspecific_fingerprints(), total=ds.nmodels, ncols=100):
|
|
337
|
+
NAM = cna.tl._nam._nam(d, sid_name, nsteps=nsteps, self_weight=self_weight,
|
|
338
|
+
show_progress=show_progress)
|
|
339
|
+
MAMs.append(NAM)
|
|
340
|
+
return MAMs
|
|
341
|
+
|
|
342
|
+
def association(ds, y, sid_name, batches=None, covs=None, donorids=None, key_added='mncoef',
|
|
343
|
+
return_full=False, ridges=None, MAMs=None,
|
|
344
|
+
Nnull=10_000, seed=0, make_umap=True,
|
|
345
|
+
nsteps=None, show_progress=False, allow_low_sample_size=False,
|
|
346
|
+
max_num_mns=5_000, **kwargs):
|
|
347
|
+
"""
|
|
348
|
+
Test patch fingerprints for association with a sample-level phenotype.
|
|
349
|
+
|
|
350
|
+
Fits a covariate-aware model relating each microniche's cross-sample
|
|
351
|
+
abundance to the phenotype, assessing significance by permutation. Returns
|
|
352
|
+
both a global p-value for the dataset and a per-microniche coefficient with
|
|
353
|
+
an empirical FDR.
|
|
354
|
+
|
|
355
|
+
Parameters
|
|
356
|
+
----------
|
|
357
|
+
ds
|
|
358
|
+
Fingerprints from `latentreps`.
|
|
359
|
+
y
|
|
360
|
+
Sample-level phenotype indexed by sample ID (e.g. case/control).
|
|
361
|
+
sid_name
|
|
362
|
+
Column in ``ds.obs`` giving each patch's sample ID.
|
|
363
|
+
batches
|
|
364
|
+
Sample-level batch labels to condition on.
|
|
365
|
+
covs
|
|
366
|
+
Sample-level covariates to control for.
|
|
367
|
+
donorids
|
|
368
|
+
Sample-level donor IDs; when given, permutations are done at the donor
|
|
369
|
+
level to respect repeated samples per donor. This cannot be used together
|
|
370
|
+
with `batches`.
|
|
371
|
+
key_added
|
|
372
|
+
Name for the per-microniche coefficient column written to ``D.obs``.
|
|
373
|
+
return_full
|
|
374
|
+
If True, return the full internal result object instead of just the
|
|
375
|
+
p-value (see Returns).
|
|
376
|
+
MAMs
|
|
377
|
+
Precomputed abundance matrices from `compute_mams`; recomputed if None.
|
|
378
|
+
Nnull
|
|
379
|
+
Number of null permutations.
|
|
380
|
+
make_umap
|
|
381
|
+
Compute a UMAP of the microniche graph for visualization.
|
|
382
|
+
max_num_mns
|
|
383
|
+
Cap on the number of microniches used for the statistical tests.
|
|
384
|
+
|
|
385
|
+
Returns
|
|
386
|
+
-------
|
|
387
|
+
tuple
|
|
388
|
+
``(p, D)`` where `p` is the global association p-value and `D` is an
|
|
389
|
+
AnnData of microniches carrying per-microniche coefficients
|
|
390
|
+
(``D.obs[key_added]``) and FDRs (``D.obs[f'{key_added}_fdr']``). If
|
|
391
|
+
`return_full` is True, returns ``(res, D)`` with the full result object
|
|
392
|
+
in place of `p`.
|
|
393
|
+
"""
|
|
394
|
+
rng = np.random.default_rng(seed)
|
|
395
|
+
np.random.seed(seed)
|
|
396
|
+
|
|
397
|
+
# Check formats of inputs and figure out which samples have valid data
|
|
398
|
+
batches, filter_samples = cna.tl._association.check_inputs(ds.select_model(0), y, sid_name, batches, covs, donorids, allow_low_sample_size)
|
|
399
|
+
|
|
400
|
+
# Compute raw NAMs (unless precomputed), then apply batch QC and sample/column filtering
|
|
401
|
+
if MAMs is None:
|
|
402
|
+
MAMs = compute_mams(ds, sid_name, nsteps=nsteps, show_progress=show_progress)
|
|
403
|
+
elif len(MAMs) != ds.nmodels:
|
|
404
|
+
raise ValueError(f'Expected MAMs of length {ds.nmodels}, got {len(MAMs)}.')
|
|
405
|
+
|
|
406
|
+
MAMs_filtered = []
|
|
407
|
+
kepts = []
|
|
408
|
+
for NAM in MAMs:
|
|
409
|
+
NAMqc, keep = cna.tl._nam._qc_nam(NAM, batches, show_progress=show_progress)
|
|
410
|
+
NAM, kept, batches, covs, donorids, filter_samples = cna.tl._association.reindex_and_filter_nam(
|
|
411
|
+
NAMqc, keep, y, batches, covs, donorids, filter_samples)
|
|
412
|
+
MAMs_filtered.append(NAM)
|
|
413
|
+
kepts.append(kept)
|
|
414
|
+
kept = np.logical_and.reduce(kepts)
|
|
415
|
+
|
|
416
|
+
for i in range(len(MAMs_filtered)):
|
|
417
|
+
MAMs_filtered[i] = MAMs_filtered[i][ds.obs.index[kept]]
|
|
418
|
+
|
|
419
|
+
# residualize NAMs
|
|
420
|
+
MAMs_concat = pd.concat(MAMs_filtered, axis=1)
|
|
421
|
+
MAMs_concat.columns = range(MAMs_concat.shape[1])
|
|
422
|
+
res = cna.tl._nam._resid_nam(MAMs_concat,
|
|
423
|
+
covs[filter_samples] if covs is not None else covs,
|
|
424
|
+
batches[filter_samples] if batches is not None else batches,
|
|
425
|
+
npcs=1,
|
|
426
|
+
ridges=ridges,
|
|
427
|
+
show_progress=show_progress)
|
|
428
|
+
MAMs_concat = res.namresid
|
|
429
|
+
|
|
430
|
+
print('performing association test')
|
|
431
|
+
n_samples, n_total = MAMs_concat.shape
|
|
432
|
+
Npatches = n_total // ds.nmodels
|
|
433
|
+
MAMresid = MAMs_concat.values.reshape(n_samples, ds.nmodels, Npatches)
|
|
434
|
+
res_ = _association(
|
|
435
|
+
MAMresid, res.M.values,
|
|
436
|
+
y[filter_samples].values, batches[filter_samples].values,
|
|
437
|
+
donorids[filter_samples].values if donorids is not None else None,
|
|
438
|
+
rng,
|
|
439
|
+
max_num_mns=max_num_mns,
|
|
440
|
+
show_progress=show_progress, Nnull=Nnull,
|
|
441
|
+
**kwargs)
|
|
442
|
+
res.__dict__.update(vars(res_)) # add info from from res_ to res
|
|
443
|
+
res.kept = kept
|
|
444
|
+
|
|
445
|
+
# make anndata with results
|
|
446
|
+
D = ds.weighted_avg_graph(res.weights, kept, make_umap=make_umap)
|
|
447
|
+
if key_added in D.obs:
|
|
448
|
+
warnings.warn(f"Key '{key_added}' already exists in d.obs. Overwriting.")
|
|
449
|
+
D.obs[key_added] = res.mncorrs
|
|
450
|
+
D.obsm['permodel_mncorrs'] = pd.DataFrame(res.permodel_mncorrs.T,
|
|
451
|
+
columns=[f'model{i}' for i in range(1, ds.nmodels+1)],
|
|
452
|
+
index=D.obs.index)
|
|
453
|
+
|
|
454
|
+
# compute local FDRs (vectorized: min fdr over all thresholds <= |mncorr|)
|
|
455
|
+
thr_sorted = res.fdrs.threshold.values # ascending, from np.quantile
|
|
456
|
+
cummin_fdr = np.minimum.accumulate(res.fdrs.fdr.values)
|
|
457
|
+
k = np.searchsorted(thr_sorted, np.abs(D.obs[key_added].values), side='right')
|
|
458
|
+
D.obs[f'{key_added}_fdr'] = np.where(k > 0, cummin_fdr[np.clip(k - 1, 0, None)], 1.0)
|
|
459
|
+
D.uns['vima_p'] = res.p
|
|
460
|
+
D.uns['vima_pheno'] = y.name
|
|
461
|
+
|
|
462
|
+
if return_full:
|
|
463
|
+
return res, D
|
|
464
|
+
else:
|
|
465
|
+
return res.p, D
|