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.
Files changed (40) hide show
  1. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/PKG-INFO +5 -2
  2. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/README.md +4 -1
  3. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/pyproject.toml +4 -0
  4. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/setup.cfg +1 -1
  5. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/__init__.py +2 -2
  6. vima_spatial-0.2.3/src/vima/cc.py +465 -0
  7. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/data/patchcollection.py +101 -0
  8. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/data/samples.py +45 -0
  9. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/fingerprints.py +88 -1
  10. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/ingest/dimreduce.py +77 -1
  11. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/ingest/ingest.py +103 -0
  12. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/ingest/nonst.py +71 -2
  13. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/ingest/st.py +150 -14
  14. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/ingest/util.py +16 -0
  15. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/models/__init__.py +23 -0
  16. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/models/simple_vae.py +3 -0
  17. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/models/vae.py +19 -0
  18. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/patchfeatures.py +105 -58
  19. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/train/logging.py +3 -0
  20. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/train/training.py +112 -1
  21. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/vis/features.py +46 -32
  22. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/vis/patches.py +129 -59
  23. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/vis/spatial.py +69 -0
  24. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/vis/umaps.py +15 -0
  25. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima_spatial.egg-info/PKG-INFO +5 -2
  26. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima_spatial.egg-info/SOURCES.txt +2 -2
  27. vima_spatial-0.2.3/tests/test_ra_regression.py +69 -0
  28. vima_spatial-0.2.2/MANIFEST.in +0 -1
  29. vima_spatial-0.2.2/src/vima/cc.py +0 -199
  30. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/data/__init__.py +0 -0
  31. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/ingest/__init__.py +0 -0
  32. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/models/resnet_vae.py +0 -0
  33. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/models/resnetlight_decoder.py +0 -0
  34. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/models/resnetlight_encoder.py +0 -0
  35. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/train/__init__.py +0 -0
  36. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/vis/__init__.py +0 -0
  37. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima/vis/patchexamples.py +0 -0
  38. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima_spatial.egg-info/dependency_links.txt +0 -0
  39. {vima_spatial-0.2.2 → vima_spatial-0.2.3}/src/vima_spatial.egg-info/requires.txt +0 -0
  40. {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.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
- Take a look at our [demo](https://github.com/yakirr/vima/blob/main/demo/demo_IF.ipynb) to see how to get started with an example analysis. We plan to put up demos for other data modalities in the future.
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
- Take a look at our [demo](https://github.com/yakirr/vima/blob/main/demo/demo_IF.ipynb) to see how to get started with an example analysis. We plan to put up demos for other data modalities in the future.
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:
@@ -4,3 +4,7 @@ requires = [
4
4
  "wheel"
5
5
  ]
6
6
  build-backend = "setuptools.build_meta"
7
+
8
+ [tool.pytest.ini_options]
9
+ testpaths = ["tests"]
10
+ addopts = "-v"
@@ -1,6 +1,6 @@
1
1
  [metadata]
2
2
  name = vima-spatial
3
- version = 0.2.2
3
+ version = 0.2.3
4
4
  author = Yakir Reshef
5
5
  author_email = yreshef@broadinstitute.org
6
6
  description = variational inference-based microniche analysis
@@ -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