vima-spatial 0.2.3__tar.gz → 0.2.4__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.3 → vima_spatial-0.2.4}/PKG-INFO +1 -1
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/setup.cfg +1 -1
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/__init__.py +5 -2
- vima_spatial-0.2.4/src/vima/_settings.py +200 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/cc.py +18 -16
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/data/__init__.py +1 -0
- vima_spatial-0.2.4/src/vima/data/download.py +62 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/data/patchcollection.py +25 -27
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/data/samples.py +2 -3
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/fingerprints.py +11 -9
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/ingest/dimreduce.py +19 -24
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/ingest/ingest.py +21 -21
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/ingest/nonst.py +9 -10
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/ingest/st.py +37 -46
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/ingest/util.py +5 -4
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/patchfeatures.py +4 -5
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/train/training.py +7 -9
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/vis/patches.py +5 -6
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/vis/spatial.py +2 -4
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/vis/umaps.py +1 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima_spatial.egg-info/PKG-INFO +1 -1
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima_spatial.egg-info/SOURCES.txt +2 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/README.md +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/pyproject.toml +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/ingest/__init__.py +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/models/__init__.py +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/models/resnet_vae.py +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/models/resnetlight_decoder.py +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/models/resnetlight_encoder.py +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/models/simple_vae.py +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/models/vae.py +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/train/__init__.py +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/train/logging.py +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/vis/__init__.py +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/vis/features.py +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/vis/patchexamples.py +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima_spatial.egg-info/dependency_links.txt +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima_spatial.egg-info/requires.txt +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima_spatial.egg-info/top_level.txt +0 -0
- {vima_spatial-0.2.3 → vima_spatial-0.2.4}/tests/test_ra_regression.py +0 -0
|
@@ -1,3 +1,4 @@
|
|
|
1
|
+
from ._settings import settings, Verbosity
|
|
1
2
|
from . import data as d
|
|
2
3
|
from . import cc
|
|
3
4
|
from . import train as t
|
|
@@ -7,12 +8,14 @@ from . import models
|
|
|
7
8
|
|
|
8
9
|
from .data.patchcollection import PatchCollection
|
|
9
10
|
from .data.samples import read_samples, reindex_by_sid
|
|
11
|
+
from .data.download import download_zenodo
|
|
10
12
|
from .train.training import train, fit, set_seed
|
|
11
13
|
from .cc import latentreps, association, compute_mams
|
|
12
14
|
from .fingerprints import Fingerprints
|
|
13
15
|
from .patchfeatures import cell_type_counts, expression_profiles, test_features
|
|
14
16
|
|
|
15
|
-
__all__ = ['
|
|
16
|
-
'
|
|
17
|
+
__all__ = ['settings', 'Verbosity',
|
|
18
|
+
'd', 'cc', 't', 'pp', 'v', 'models',
|
|
19
|
+
'PatchCollection', 'read_samples', 'reindex_by_sid', 'download_zenodo',
|
|
17
20
|
'train', 'fit', 'set_seed', 'latentreps', 'association', 'compute_mams',
|
|
18
21
|
'Fingerprints', 'cc', 'cell_type_counts', 'expression_profiles', 'test_features']
|
|
@@ -0,0 +1,200 @@
|
|
|
1
|
+
"""Unified output control for vima.
|
|
2
|
+
|
|
3
|
+
All informational output in the package flows through a single ``logging``
|
|
4
|
+
logger (``vima._settings.logger``, named ``"vima"``) and a single progress-bar
|
|
5
|
+
helper (``settings.progress``). User-facing controls live on the module-level
|
|
6
|
+
``settings`` object, exposed as ``vima.settings``.
|
|
7
|
+
|
|
8
|
+
Three independent knobs:
|
|
9
|
+
|
|
10
|
+
* ``settings.verbosity`` -- how much informational output to emit. Three levels:
|
|
11
|
+
|
|
12
|
+
================== === ================ ==================================
|
|
13
|
+
name int logging level what is shown
|
|
14
|
+
================== === ================ ==================================
|
|
15
|
+
``"minimal"`` 0 ``WARNING`` warnings/errors, progress bars,
|
|
16
|
+
and results only (e.g. the
|
|
17
|
+
``association`` p-value)
|
|
18
|
+
``"default"`` 1 ``INFO`` + high-level progress messages
|
|
19
|
+
``"verbose"`` 2 ``DEBUG`` + detailed diagnostics
|
|
20
|
+
================== === ================ ==================================
|
|
21
|
+
|
|
22
|
+
* ``settings.diagnostic_plots`` -- how many diagnostic plots to draw, on the same
|
|
23
|
+
three-level scale as ``verbosity``: ``"minimal"`` draws none, ``"default"`` draws
|
|
24
|
+
the plots historically shown by default (QC, mean-variance, embeddings), and
|
|
25
|
+
``"verbose"`` additionally draws the more detailed plots. Until it is set
|
|
26
|
+
explicitly, it *tracks* ``verbosity`` -- so setting ``verbosity`` alone also
|
|
27
|
+
controls the plots; assigning ``diagnostic_plots`` decouples the two.
|
|
28
|
+
|
|
29
|
+
* ``settings.progress_bars`` -- whether tqdm progress bars are shown. Orthogonal
|
|
30
|
+
to verbosity (bars stay on even at ``"minimal"``); set ``False`` for batch or
|
|
31
|
+
cluster runs.
|
|
32
|
+
|
|
33
|
+
Guidelines for package code:
|
|
34
|
+
|
|
35
|
+
* ``logger.info(...)`` -- normal progress messages (visible at ``default``).
|
|
36
|
+
* ``logger.debug(...)`` -- detailed diagnostics (visible only at ``verbose``).
|
|
37
|
+
* ``logger.warning(...)`` -- warnings the user should see at every level.
|
|
38
|
+
* ``settings.result(...)`` -- user-facing results shown at every level.
|
|
39
|
+
* ``settings.progress(iterable, name=...)`` -- wrap any loop needing a bar.
|
|
40
|
+
* ``if settings.show_plots(...):`` -- guard a diagnostic plot. Pass ``"verbose"``
|
|
41
|
+
for the detailed plots; the default level guards the standard ones.
|
|
42
|
+
"""
|
|
43
|
+
|
|
44
|
+
import sys
|
|
45
|
+
import logging
|
|
46
|
+
from enum import IntEnum
|
|
47
|
+
|
|
48
|
+
from tqdm.auto import tqdm
|
|
49
|
+
|
|
50
|
+
__all__ = ["Verbosity", "settings", "logger"]
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class Verbosity(IntEnum):
|
|
54
|
+
"""Amount of informational output emitted by vima."""
|
|
55
|
+
|
|
56
|
+
minimal = 0 # warnings + progress bars + results only
|
|
57
|
+
default = 1 # + high-level info messages (the historical default)
|
|
58
|
+
verbose = 2 # + detailed diagnostics
|
|
59
|
+
|
|
60
|
+
@classmethod
|
|
61
|
+
def parse(cls, value):
|
|
62
|
+
"""Coerce an int, name string, or Verbosity into a Verbosity."""
|
|
63
|
+
if isinstance(value, cls):
|
|
64
|
+
return value
|
|
65
|
+
if isinstance(value, str):
|
|
66
|
+
try:
|
|
67
|
+
return cls[value]
|
|
68
|
+
except KeyError:
|
|
69
|
+
raise ValueError(
|
|
70
|
+
f"unknown verbosity {value!r}; "
|
|
71
|
+
f"expected one of {[v.name for v in cls]} or 0/1/2"
|
|
72
|
+
)
|
|
73
|
+
return cls(value)
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
_LOGGING_LEVELS = {
|
|
77
|
+
Verbosity.minimal: logging.WARNING,
|
|
78
|
+
Verbosity.default: logging.INFO,
|
|
79
|
+
Verbosity.verbose: logging.DEBUG,
|
|
80
|
+
}
|
|
81
|
+
|
|
82
|
+
logger = logging.getLogger("vima")
|
|
83
|
+
|
|
84
|
+
# ANSI colors: debug messages are grayed out, results are green.
|
|
85
|
+
_GRAY, _GREEN, _RESET = "\033[90m", "\033[32m", "\033[0m"
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
class _ColorFormatter(logging.Formatter):
|
|
89
|
+
"""Format log records, graying out DEBUG-level messages."""
|
|
90
|
+
|
|
91
|
+
def format(self, record):
|
|
92
|
+
msg = super().format(record)
|
|
93
|
+
if record.levelno == logging.DEBUG:
|
|
94
|
+
return f"{_GRAY}{msg}{_RESET}"
|
|
95
|
+
return msg
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
class _TqdmLoggingHandler(logging.StreamHandler):
|
|
99
|
+
"""Route log records through ``tqdm.write`` so they never corrupt an active
|
|
100
|
+
progress bar. When no bar is live, ``tqdm.write`` degrades to a plain write
|
|
101
|
+
to the stream, so logging in loops without a bar is unaffected. This keeps
|
|
102
|
+
the two output channels fully decoupled: progress bars honor
|
|
103
|
+
``settings.progress_bars`` and log messages honor ``settings.verbosity``,
|
|
104
|
+
independent of each other.
|
|
105
|
+
"""
|
|
106
|
+
|
|
107
|
+
def emit(self, record):
|
|
108
|
+
try:
|
|
109
|
+
msg = self.format(record)
|
|
110
|
+
tqdm.write(msg, file=self.stream, end=self.terminator)
|
|
111
|
+
except Exception:
|
|
112
|
+
self.handleError(record)
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
class Settings:
|
|
116
|
+
"""Global vima output settings; access via the ``vima.settings`` singleton."""
|
|
117
|
+
|
|
118
|
+
def __init__(self):
|
|
119
|
+
# Attach a single stdout handler and keep our records off the root
|
|
120
|
+
# logger so vima's output is independent of any host logging config.
|
|
121
|
+
# The "vima" logger is a process-global singleton that survives module
|
|
122
|
+
# reloads, so clear any handler we previously attached before adding a
|
|
123
|
+
# new one; otherwise re-importing vima (e.g. importlib.reload in a
|
|
124
|
+
# notebook) accumulates handlers and every message prints N times.
|
|
125
|
+
# Match by class *name*, not isinstance: each reload defines a fresh
|
|
126
|
+
# _TqdmLoggingHandler class, so handlers left by earlier reloads are
|
|
127
|
+
# instances of a different (stale) class object and would fail an
|
|
128
|
+
# isinstance check against the current class.
|
|
129
|
+
for h in list(logger.handlers):
|
|
130
|
+
if type(h).__name__ == "_TqdmLoggingHandler":
|
|
131
|
+
logger.removeHandler(h)
|
|
132
|
+
handler = _TqdmLoggingHandler(sys.stdout)
|
|
133
|
+
handler.setFormatter(_ColorFormatter("%(message)s"))
|
|
134
|
+
logger.addHandler(handler)
|
|
135
|
+
logger.propagate = False
|
|
136
|
+
|
|
137
|
+
self.progress_bars = True
|
|
138
|
+
self._verbosity = None
|
|
139
|
+
# ``diagnostic_plots`` tracks ``verbosity`` until the user sets it.
|
|
140
|
+
self._diagnostic_plots = None
|
|
141
|
+
self._diagnostic_plots_explicit = False
|
|
142
|
+
self.verbosity = Verbosity.default
|
|
143
|
+
|
|
144
|
+
@property
|
|
145
|
+
def verbosity(self):
|
|
146
|
+
"""Current :class:`Verbosity` level (see module docstring)."""
|
|
147
|
+
return self._verbosity
|
|
148
|
+
|
|
149
|
+
@verbosity.setter
|
|
150
|
+
def verbosity(self, value):
|
|
151
|
+
self._verbosity = Verbosity.parse(value)
|
|
152
|
+
logger.setLevel(_LOGGING_LEVELS[self._verbosity])
|
|
153
|
+
if not self._diagnostic_plots_explicit:
|
|
154
|
+
self._diagnostic_plots = self._verbosity
|
|
155
|
+
|
|
156
|
+
@property
|
|
157
|
+
def diagnostic_plots(self):
|
|
158
|
+
""":class:`Verbosity` level controlling how many diagnostic plots are drawn.
|
|
159
|
+
|
|
160
|
+
Tracks :attr:`verbosity` until assigned; assigning it decouples the two.
|
|
161
|
+
"""
|
|
162
|
+
return self._diagnostic_plots
|
|
163
|
+
|
|
164
|
+
@diagnostic_plots.setter
|
|
165
|
+
def diagnostic_plots(self, value):
|
|
166
|
+
self._diagnostic_plots = Verbosity.parse(value)
|
|
167
|
+
self._diagnostic_plots_explicit = True
|
|
168
|
+
|
|
169
|
+
def show_plots(self, level=Verbosity.default):
|
|
170
|
+
"""Whether diagnostic plots at ``level`` should be drawn.
|
|
171
|
+
|
|
172
|
+
``level`` is anything :meth:`Verbosity.parse` accepts (default the
|
|
173
|
+
standard-plot level); pass ``"verbose"`` for the more detailed plots.
|
|
174
|
+
"""
|
|
175
|
+
return self._diagnostic_plots >= Verbosity.parse(level)
|
|
176
|
+
|
|
177
|
+
def progress(self, iterable=None, name=None, total=None, ncols=100, desc=None, **kwargs):
|
|
178
|
+
"""tqdm wrapper honoring ``settings.progress_bars``.
|
|
179
|
+
|
|
180
|
+
Central replacement for the per-module ``pb = lambda ...`` helpers.
|
|
181
|
+
``name`` is a brief label for the loop, shown as the tqdm ``desc``
|
|
182
|
+
prefix (``desc`` is still accepted as an alias; ``name`` wins).
|
|
183
|
+
"""
|
|
184
|
+
using_widget = any(
|
|
185
|
+
cls.__module__ == "tqdm.notebook"
|
|
186
|
+
for cls in tqdm.mro()
|
|
187
|
+
)
|
|
188
|
+
return tqdm(
|
|
189
|
+
iterable, desc=name if name is not None else desc, total=total,
|
|
190
|
+
ncols=ncols * (7 if using_widget else 1),
|
|
191
|
+
bar_format='{l_bar}{bar}{r_bar}',
|
|
192
|
+
disable=not self.progress_bars, **kwargs,
|
|
193
|
+
)
|
|
194
|
+
|
|
195
|
+
def result(self, message):
|
|
196
|
+
"""Emit a user-facing result, shown at every verbosity level."""
|
|
197
|
+
tqdm.write(f"{_GREEN}{message}{_RESET}", file=sys.stdout)
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
settings = Settings()
|
|
@@ -1,4 +1,4 @@
|
|
|
1
|
-
import os
|
|
1
|
+
import os, gc
|
|
2
2
|
from concurrent.futures import ThreadPoolExecutor
|
|
3
3
|
import numpy as np
|
|
4
4
|
import scanpy as sc
|
|
@@ -11,9 +11,8 @@ import warnings
|
|
|
11
11
|
import scipy.sparse as sp
|
|
12
12
|
import scipy.stats as st
|
|
13
13
|
from argparse import Namespace
|
|
14
|
-
from tqdm import tqdm
|
|
15
14
|
from .fingerprints import Fingerprints
|
|
16
|
-
|
|
15
|
+
from ._settings import Verbosity, settings, logger
|
|
17
16
|
|
|
18
17
|
def anndata(patchmeta, Z, var_names=None, use_rep='X', n_comps=10, **kwargs):
|
|
19
18
|
"""
|
|
@@ -86,7 +85,7 @@ def apply(models, P, batch_size=1000, with_mse=False):
|
|
|
86
85
|
if with_mse:
|
|
87
86
|
MSEs = {modelid: [] for modelid in range(len(models))}
|
|
88
87
|
with torch.no_grad():
|
|
89
|
-
for batch in
|
|
88
|
+
for batch in settings.progress(eval_loader, name='apply models'):
|
|
90
89
|
for modelid, model in enumerate(models):
|
|
91
90
|
if with_mse:
|
|
92
91
|
x_recon, mean, _ = model(batch, sample_from_latent=False)
|
|
@@ -129,12 +128,12 @@ def latentreps(models, P, use_rep='X', n_comps=100, with_mse=True, **kwargs):
|
|
|
129
128
|
Per-model embeddings and neighbor graphs, with ``.obs`` carrying patch
|
|
130
129
|
metadata.
|
|
131
130
|
"""
|
|
132
|
-
|
|
131
|
+
logger.info('applying models')
|
|
133
132
|
result = apply(models, P, with_mse=with_mse)
|
|
134
133
|
Zs, MSEs = result if with_mse else (result, None)
|
|
135
134
|
|
|
136
|
-
|
|
137
|
-
ds = [anndata(P.meta, Z, use_rep=use_rep, n_comps=n_comps, **kwargs) for Z in
|
|
135
|
+
logger.info('computing nearest-neighbor graphs')
|
|
136
|
+
ds = [anndata(P.meta, Z, use_rep=use_rep, n_comps=n_comps, **kwargs) for Z in settings.progress(Zs, name='nearest-neighbor graphs')]
|
|
138
137
|
fp = Fingerprints.from_list(ds)
|
|
139
138
|
|
|
140
139
|
if MSEs is not None:
|
|
@@ -288,7 +287,7 @@ def _association(MAMresid, M, y, batches, donorids, rng, Nnull=10_000,
|
|
|
288
287
|
|
|
289
288
|
# compute global p-vaule
|
|
290
289
|
p = ((nullglobalstats >= globalstat).sum() + 1)/(len(nullglobalstats) + 1)
|
|
291
|
-
|
|
290
|
+
settings.result(f'\033[32mP = {p}\033[0m')
|
|
292
291
|
if p <= 1/(Nnull + 1)+1e-10:
|
|
293
292
|
warnings.warn('global association p-value attained minimal possible value. '+\
|
|
294
293
|
'Consider increasing Nnull')
|
|
@@ -302,16 +301,14 @@ def _association(MAMresid, M, y, batches, donorids, rng, Nnull=10_000,
|
|
|
302
301
|
res = {'p':p, 'mncorrs':mncorrs_meta, 'fdrs':fdrs,
|
|
303
302
|
'globalstat':globalstat, 'nullglobalstats':nullglobalstats,
|
|
304
303
|
'weights':weights,
|
|
305
|
-
'nullmncorrs':nullmncorrs_meta,
|
|
306
304
|
'permodel_mncorrs':mncorrs,
|
|
307
|
-
'MAMres':MAMresid,
|
|
308
305
|
'ycond':ycond
|
|
309
306
|
}
|
|
310
307
|
|
|
311
308
|
return Namespace(**res)
|
|
312
309
|
|
|
313
310
|
|
|
314
|
-
def compute_mams(ds, sid_name, nsteps=None, self_weight=1, show_progress=
|
|
311
|
+
def compute_mams(ds, sid_name, nsteps=None, self_weight=1, show_progress=None):
|
|
315
312
|
"""
|
|
316
313
|
Precompute per-model sample-by-microniche abundance matrices.
|
|
317
314
|
|
|
@@ -331,9 +328,11 @@ def compute_mams(ds, sid_name, nsteps=None, self_weight=1, show_progress=False):
|
|
|
331
328
|
list
|
|
332
329
|
One abundance matrix per model.
|
|
333
330
|
"""
|
|
334
|
-
|
|
331
|
+
if show_progress is None:
|
|
332
|
+
show_progress = settings.verbosity >= Verbosity.verbose
|
|
333
|
+
logger.info('computing MAT') #TODO: rename MAM to MAT in code if we keep this nomenclature
|
|
335
334
|
MAMs = []
|
|
336
|
-
for d in
|
|
335
|
+
for d in settings.progress(ds.modelspecific_fingerprints(), total=ds.nmodels, name='compute MAMs'):
|
|
337
336
|
NAM = cna.tl._nam._nam(d, sid_name, nsteps=nsteps, self_weight=self_weight,
|
|
338
337
|
show_progress=show_progress)
|
|
339
338
|
MAMs.append(NAM)
|
|
@@ -342,7 +341,7 @@ def compute_mams(ds, sid_name, nsteps=None, self_weight=1, show_progress=False):
|
|
|
342
341
|
def association(ds, y, sid_name, batches=None, covs=None, donorids=None, key_added='mncoef',
|
|
343
342
|
return_full=False, ridges=None, MAMs=None,
|
|
344
343
|
Nnull=10_000, seed=0, make_umap=True,
|
|
345
|
-
nsteps=None, show_progress=
|
|
344
|
+
nsteps=None, show_progress=None, allow_low_sample_size=False,
|
|
346
345
|
max_num_mns=5_000, **kwargs):
|
|
347
346
|
"""
|
|
348
347
|
Test patch fingerprints for association with a sample-level phenotype.
|
|
@@ -391,6 +390,8 @@ def association(ds, y, sid_name, batches=None, covs=None, donorids=None, key_add
|
|
|
391
390
|
`return_full` is True, returns ``(res, D)`` with the full result object
|
|
392
391
|
in place of `p`.
|
|
393
392
|
"""
|
|
393
|
+
if show_progress is None:
|
|
394
|
+
show_progress = settings.verbosity >= Verbosity.verbose
|
|
394
395
|
rng = np.random.default_rng(seed)
|
|
395
396
|
np.random.seed(seed)
|
|
396
397
|
|
|
@@ -414,7 +415,7 @@ def association(ds, y, sid_name, batches=None, covs=None, donorids=None, key_add
|
|
|
414
415
|
kept = np.logical_and.reduce(kepts)
|
|
415
416
|
|
|
416
417
|
for i in range(len(MAMs_filtered)):
|
|
417
|
-
MAMs_filtered[i] = MAMs_filtered[i]
|
|
418
|
+
MAMs_filtered[i] = MAMs_filtered[i].iloc[:,kept]
|
|
418
419
|
|
|
419
420
|
# residualize NAMs
|
|
420
421
|
MAMs_concat = pd.concat(MAMs_filtered, axis=1)
|
|
@@ -427,7 +428,7 @@ def association(ds, y, sid_name, batches=None, covs=None, donorids=None, key_add
|
|
|
427
428
|
show_progress=show_progress)
|
|
428
429
|
MAMs_concat = res.namresid
|
|
429
430
|
|
|
430
|
-
|
|
431
|
+
logger.info('performing association test')
|
|
431
432
|
n_samples, n_total = MAMs_concat.shape
|
|
432
433
|
Npatches = n_total // ds.nmodels
|
|
433
434
|
MAMresid = MAMs_concat.values.reshape(n_samples, ds.nmodels, Npatches)
|
|
@@ -441,6 +442,7 @@ def association(ds, y, sid_name, batches=None, covs=None, donorids=None, key_add
|
|
|
441
442
|
**kwargs)
|
|
442
443
|
res.__dict__.update(vars(res_)) # add info from from res_ to res
|
|
443
444
|
res.kept = kept
|
|
445
|
+
gc.collect()
|
|
444
446
|
|
|
445
447
|
# make anndata with results
|
|
446
448
|
D = ds.weighted_avg_graph(res.weights, kept, make_umap=make_umap)
|
|
@@ -0,0 +1,62 @@
|
|
|
1
|
+
"""Utilities for fetching example/demo datasets."""
|
|
2
|
+
|
|
3
|
+
import os
|
|
4
|
+
import json
|
|
5
|
+
from urllib.request import urlopen, Request
|
|
6
|
+
|
|
7
|
+
from .._settings import settings, logger
|
|
8
|
+
|
|
9
|
+
__all__ = ["download_zenodo"]
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def download_zenodo(record_id, target_dir, chunk_size=1024 * 1024):
|
|
13
|
+
"""Download every file attached to a Zenodo record into ``target_dir``.
|
|
14
|
+
|
|
15
|
+
Parameters
|
|
16
|
+
----------
|
|
17
|
+
record_id
|
|
18
|
+
Zenodo record id -- the number in the record URL, e.g. ``"21535534"``.
|
|
19
|
+
target_dir
|
|
20
|
+
Directory to write the files into; created if it does not exist.
|
|
21
|
+
chunk_size
|
|
22
|
+
Streaming chunk size in bytes.
|
|
23
|
+
"""
|
|
24
|
+
os.makedirs(target_dir, exist_ok=True)
|
|
25
|
+
with urlopen(f"https://zenodo.org/api/records/{record_id}") as r:
|
|
26
|
+
files = json.load(r)["files"]
|
|
27
|
+
for i, f in enumerate(files):
|
|
28
|
+
url = f["links"]["self"]
|
|
29
|
+
fname = f["key"]
|
|
30
|
+
size = f.get("size", 0)
|
|
31
|
+
dest = os.path.join(target_dir, fname)
|
|
32
|
+
logger.info(f"[{i + 1}/{len(files)}] downloading {fname}")
|
|
33
|
+
req = Request(url, headers={"User-Agent": "vima"})
|
|
34
|
+
with urlopen(req) as resp, open(dest, "wb") as out, \
|
|
35
|
+
settings.progress(total=size, name=fname,
|
|
36
|
+
unit="B", unit_scale=True) as pbar:
|
|
37
|
+
while True:
|
|
38
|
+
chunk = resp.read(chunk_size)
|
|
39
|
+
if not chunk:
|
|
40
|
+
break
|
|
41
|
+
out.write(chunk)
|
|
42
|
+
pbar.update(len(chunk))
|
|
43
|
+
|
|
44
|
+
def download_toy_rawdata(target_dir):
|
|
45
|
+
"""Download the toy raw data for the demo.
|
|
46
|
+
|
|
47
|
+
Parameters
|
|
48
|
+
----------
|
|
49
|
+
target_dir
|
|
50
|
+
Directory to write the files into; created if it does not exist.
|
|
51
|
+
"""
|
|
52
|
+
download_zenodo("20433752", target_dir)
|
|
53
|
+
|
|
54
|
+
def download_toy_metadata_and_fingerprints(target_dir):
|
|
55
|
+
"""Download the toy metadata and precomputed fingerprints for the demo.
|
|
56
|
+
|
|
57
|
+
Parameters
|
|
58
|
+
----------
|
|
59
|
+
target_dir
|
|
60
|
+
Directory to write the files into; created if it does not exist.
|
|
61
|
+
"""
|
|
62
|
+
download_zenodo("21535534", target_dir)
|
|
@@ -4,10 +4,10 @@ from torchvision import transforms
|
|
|
4
4
|
import numpy as np
|
|
5
5
|
import pandas as pd
|
|
6
6
|
import random
|
|
7
|
+
import logging
|
|
7
8
|
import torch
|
|
8
9
|
from . import samples as vds
|
|
9
|
-
from
|
|
10
|
-
pb = lambda x: tqdm(x, ncols=100)
|
|
10
|
+
from .._settings import settings, logger
|
|
11
11
|
|
|
12
12
|
class ToTorch:
|
|
13
13
|
"""Transform converting an array to a channels-first torch tensor."""
|
|
@@ -55,7 +55,7 @@ class PatchCollection(Dataset):
|
|
|
55
55
|
"""
|
|
56
56
|
|
|
57
57
|
@staticmethod
|
|
58
|
-
def choose_patches(samples, patchsize, patchstride, max_frac_empty
|
|
58
|
+
def choose_patches(samples, patchsize, patchstride, max_frac_empty):
|
|
59
59
|
"""
|
|
60
60
|
Pick patch grid positions for each sample, keeping only patches with
|
|
61
61
|
enough non-empty pixels.
|
|
@@ -72,7 +72,7 @@ class PatchCollection(Dataset):
|
|
|
72
72
|
"""
|
|
73
73
|
patchmeta = []
|
|
74
74
|
|
|
75
|
-
for s in
|
|
75
|
+
for s in settings.progress(samples.values(), name='choose patches'):
|
|
76
76
|
mask = vds.get_mask(s)
|
|
77
77
|
starts = np.array([
|
|
78
78
|
[i, j]
|
|
@@ -123,7 +123,7 @@ class PatchCollection(Dataset):
|
|
|
123
123
|
self.meta[col] = pd.factorize(self.meta.sid.map(mapping))[0]
|
|
124
124
|
self._covariate_cols.append(col)
|
|
125
125
|
|
|
126
|
-
def compute_stats(self, percentile_thresh
|
|
126
|
+
def compute_stats(self, percentile_thresh):
|
|
127
127
|
"""
|
|
128
128
|
Compute per-marker mean, std, and display percentiles over a random
|
|
129
129
|
subset of patches.
|
|
@@ -139,12 +139,12 @@ class PatchCollection(Dataset):
|
|
|
139
139
|
self.vmin = (-self.means - self.percentiles)/self.stds
|
|
140
140
|
self.vmax = (-self.means + self.percentiles)/self.stds
|
|
141
141
|
|
|
142
|
-
if
|
|
142
|
+
if logger.isEnabledFor(logging.DEBUG):
|
|
143
143
|
fmt = lambda a: ' '.join(f'{v:.2g}' for v in a)
|
|
144
|
-
|
|
145
|
-
|
|
144
|
+
logger.debug(f'per-channel means: {fmt(self.means)}')
|
|
145
|
+
logger.debug(f'per-channel stds: {fmt(self.stds)}')
|
|
146
146
|
|
|
147
|
-
def normalize(self, normalization
|
|
147
|
+
def normalize(self, normalization):
|
|
148
148
|
"""
|
|
149
149
|
Apply the chosen per-marker normalization to the stored patches.
|
|
150
150
|
|
|
@@ -157,8 +157,8 @@ class PatchCollection(Dataset):
|
|
|
157
157
|
if normalization is not None and normalization not in ['center', 'standardize', 'none']:
|
|
158
158
|
raise ValueError('normalization must equal "standardize" | "center" | "none" | None')
|
|
159
159
|
|
|
160
|
-
|
|
161
|
-
|
|
160
|
+
logger.debug(f'Normalizing color channels (normalization={normalization})...')
|
|
161
|
+
|
|
162
162
|
self.empty = np.zeros(self.patches.shape[-1], dtype=np.float32)
|
|
163
163
|
if normalization == 'standardize' or normalization == 'center':
|
|
164
164
|
self.patches = self.patches - self.means[None,None,None,:]
|
|
@@ -169,22 +169,21 @@ class PatchCollection(Dataset):
|
|
|
169
169
|
self.normalization = normalization
|
|
170
170
|
|
|
171
171
|
def __init__(self, samples, patchsize=40, patchstride=10, max_frac_empty=0.8,
|
|
172
|
-
normalization='standardize', percentile_thresh=99,
|
|
172
|
+
normalization='standardize', percentile_thresh=99,
|
|
173
173
|
covariates=None, condition_on_sid=True):
|
|
174
174
|
self.samples = samples
|
|
175
175
|
self.patchstride = patchstride
|
|
176
176
|
self._covariate_cols = []
|
|
177
|
-
self.meta = PatchCollection.choose_patches(samples, patchsize, patchstride, max_frac_empty
|
|
177
|
+
self.meta = PatchCollection.choose_patches(samples, patchsize, patchstride, max_frac_empty)
|
|
178
178
|
self.nmarkers = next(iter(samples.values())).sizes['marker']
|
|
179
179
|
|
|
180
180
|
self.pytorch_mode()
|
|
181
181
|
self.make_patchmeta(covariates=covariates, condition_on_sid=condition_on_sid)
|
|
182
|
-
self.compute_stats(percentile_thresh
|
|
183
|
-
self.normalize(normalization=normalization
|
|
182
|
+
self.compute_stats(percentile_thresh)
|
|
183
|
+
self.normalize(normalization=normalization)
|
|
184
184
|
self.augmentation_off()
|
|
185
185
|
|
|
186
|
-
def refined(self, max_frac_empty, tol=1e-10, normalization='standardize', percentile_thresh=99
|
|
187
|
-
verbose=False):
|
|
186
|
+
def refined(self, max_frac_empty, tol=1e-10, normalization='standardize', percentile_thresh=99):
|
|
188
187
|
"""
|
|
189
188
|
Return a copy restricted to denser patches.
|
|
190
189
|
|
|
@@ -202,8 +201,7 @@ class PatchCollection(Dataset):
|
|
|
202
201
|
keep = np.where(empty_frac < max_frac_empty)[0]
|
|
203
202
|
result = copy.copy(self)
|
|
204
203
|
result.subset(keep,
|
|
205
|
-
normalization=normalization, percentile_thresh=percentile_thresh
|
|
206
|
-
verbose=verbose)
|
|
204
|
+
normalization=normalization, percentile_thresh=percentile_thresh)
|
|
207
205
|
return result
|
|
208
206
|
|
|
209
207
|
@property
|
|
@@ -224,9 +222,9 @@ class PatchCollection(Dataset):
|
|
|
224
222
|
def augmentation_on(self):
|
|
225
223
|
"""Enable random rotation and horizontal flip augmentation (pytorch mode only)."""
|
|
226
224
|
if self.dim_order != 'pytorch':
|
|
227
|
-
|
|
225
|
+
logger.warning('Data augmentation only available in pytorch mode. Will leave augmentation off')
|
|
228
226
|
return
|
|
229
|
-
|
|
227
|
+
logger.debug('[PatchCollection: data augmentation on]')
|
|
230
228
|
self.transform = transforms.Compose([
|
|
231
229
|
ToTorch(),
|
|
232
230
|
RandomDiscreteRotation(),
|
|
@@ -234,7 +232,7 @@ class PatchCollection(Dataset):
|
|
|
234
232
|
])
|
|
235
233
|
def augmentation_off(self):
|
|
236
234
|
"""Disable rotation and flip augmentation."""
|
|
237
|
-
|
|
235
|
+
logger.debug('[PatchCollection: data augmentation is off]')
|
|
238
236
|
self.transform = transforms.Compose([
|
|
239
237
|
ToTorch(),
|
|
240
238
|
])
|
|
@@ -246,22 +244,22 @@ class PatchCollection(Dataset):
|
|
|
246
244
|
def pytorch_mode(self):
|
|
247
245
|
"""Switch patch output to channels-first (C, H, W) torch layout."""
|
|
248
246
|
self.dim_order = 'pytorch'
|
|
249
|
-
|
|
247
|
+
logger.debug('[PatchCollection: in pytorch mode]')
|
|
250
248
|
def numpy_mode(self):
|
|
251
249
|
"""Switch patch output to ``(H, W, C)`` numpy layout and turn augmentation off."""
|
|
252
250
|
self.dim_order = 'numpy'
|
|
253
251
|
self.augmentation_off()
|
|
254
|
-
|
|
252
|
+
logger.debug('[PatchCollection: in numpy mode]')
|
|
255
253
|
|
|
256
|
-
def subset(self, ix, percentile_thresh, normalization
|
|
254
|
+
def subset(self, ix, percentile_thresh, normalization):
|
|
257
255
|
"""
|
|
258
256
|
Restrict the collection in place to the given patch indices and recompute
|
|
259
257
|
normalization statistics.
|
|
260
258
|
"""
|
|
261
259
|
self.patches = self.patches[ix]
|
|
262
260
|
self.meta = self.meta.iloc[ix]
|
|
263
|
-
self.compute_stats(percentile_thresh
|
|
264
|
-
self.normalize(normalization=normalization
|
|
261
|
+
self.compute_stats(percentile_thresh)
|
|
262
|
+
self.normalize(normalization=normalization)
|
|
265
263
|
|
|
266
264
|
def __repr__(self):
|
|
267
265
|
ps = self.meta.patchsize.iloc[0]
|
|
@@ -5,8 +5,7 @@ import xarray as xr
|
|
|
5
5
|
import pandas as pd
|
|
6
6
|
import matplotlib.pyplot as plt
|
|
7
7
|
import cv2 as cv2
|
|
8
|
-
from
|
|
9
|
-
pb = lambda x: tqdm(x, ncols=100)
|
|
8
|
+
from .._settings import settings
|
|
10
9
|
|
|
11
10
|
def read_samples(files, stop_after=None):
|
|
12
11
|
"""
|
|
@@ -30,7 +29,7 @@ def read_samples(files, stop_after=None):
|
|
|
30
29
|
if stop_after is None: stop_after = len(files)
|
|
31
30
|
|
|
32
31
|
samples = {}
|
|
33
|
-
for f in
|
|
32
|
+
for f in settings.progress(files[:stop_after], name='read samples'):
|
|
34
33
|
s = xr.open_dataarray(f).astype(np.float32)
|
|
35
34
|
s.attrs['sid'] = os.path.splitext(os.path.basename(f))[0]
|
|
36
35
|
samples[s.sid] = s
|
|
@@ -4,9 +4,7 @@ import anndata as ad
|
|
|
4
4
|
import scanpy as sc
|
|
5
5
|
import scipy.sparse as sp
|
|
6
6
|
import cna
|
|
7
|
-
import
|
|
8
|
-
from tqdm import tqdm
|
|
9
|
-
pb = lambda x: tqdm(x, ncols=100)
|
|
7
|
+
from ._settings import settings, logger
|
|
10
8
|
|
|
11
9
|
|
|
12
10
|
class _ObsmView:
|
|
@@ -79,11 +77,15 @@ class Fingerprints:
|
|
|
79
77
|
return Fingerprints(self._adata[key].copy())
|
|
80
78
|
|
|
81
79
|
def select_model(self, i):
|
|
82
|
-
"""
|
|
83
|
-
|
|
80
|
+
"""
|
|
81
|
+
Return model `i`'s embedding and neighbor graph as an AnnData. Does not
|
|
82
|
+
make copies of the underlying data, so modifications may affect the
|
|
83
|
+
original Fingerprints object as well.
|
|
84
|
+
"""
|
|
85
|
+
d = ad.AnnData(X=self._adata.obsm[f'X_{i}'], obs=self._adata.obs)
|
|
84
86
|
d.obsp['connectivities'] = self._adata.obsp[f'connectivities_{i}']
|
|
85
87
|
d.obsp['distances'] = self._adata.obsp[f'distances_{i}']
|
|
86
|
-
d.obsm =
|
|
88
|
+
d.obsm = dict(self.obsm.items())
|
|
87
89
|
d.uns['neighbors'] = self._adata.uns[f'neighbors_{i}']
|
|
88
90
|
return d
|
|
89
91
|
|
|
@@ -151,7 +153,7 @@ class Fingerprints:
|
|
|
151
153
|
'n_neighbors': 15, 'use_rep': 'X', 'n_pcs': None},
|
|
152
154
|
}
|
|
153
155
|
if make_umap:
|
|
154
|
-
|
|
156
|
+
logger.info('Computing UMAP...')
|
|
155
157
|
sc.tl.umap(D, neighbors_key='neighbors')
|
|
156
158
|
return D
|
|
157
159
|
|
|
@@ -165,7 +167,7 @@ class Fingerprints:
|
|
|
165
167
|
|
|
166
168
|
def compute_nngs(self, **kwargs):
|
|
167
169
|
"""Recompute and store each model's nearest-neighbor graph."""
|
|
168
|
-
for i in
|
|
170
|
+
for i in settings.progress(range(self.nmodels)):
|
|
169
171
|
d = self.select_model(i)
|
|
170
172
|
sc.pp.neighbors(d, **kwargs)
|
|
171
173
|
self._adata.obsp[f'connectivities_{i}'] = d.obsp['connectivities']
|
|
@@ -220,7 +222,7 @@ class Fingerprints:
|
|
|
220
222
|
NAM -= NAM.mean(axis=0)
|
|
221
223
|
NAM /= NAM.std(axis=0)
|
|
222
224
|
_, _, VT = np.linalg.svd(NAM, full_matrices=False)
|
|
223
|
-
|
|
225
|
+
logger.debug('VT shape: %s', VT.shape)
|
|
224
226
|
return pd.DataFrame(VT.T, index=NAM.columns,
|
|
225
227
|
columns=[f'PC{i+1}' for i in range(VT.shape[0])])
|
|
226
228
|
|