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.
Files changed (40) hide show
  1. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/PKG-INFO +1 -1
  2. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/setup.cfg +1 -1
  3. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/__init__.py +5 -2
  4. vima_spatial-0.2.4/src/vima/_settings.py +200 -0
  5. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/cc.py +18 -16
  6. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/data/__init__.py +1 -0
  7. vima_spatial-0.2.4/src/vima/data/download.py +62 -0
  8. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/data/patchcollection.py +25 -27
  9. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/data/samples.py +2 -3
  10. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/fingerprints.py +11 -9
  11. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/ingest/dimreduce.py +19 -24
  12. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/ingest/ingest.py +21 -21
  13. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/ingest/nonst.py +9 -10
  14. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/ingest/st.py +37 -46
  15. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/ingest/util.py +5 -4
  16. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/patchfeatures.py +4 -5
  17. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/train/training.py +7 -9
  18. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/vis/patches.py +5 -6
  19. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/vis/spatial.py +2 -4
  20. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/vis/umaps.py +1 -0
  21. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima_spatial.egg-info/PKG-INFO +1 -1
  22. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima_spatial.egg-info/SOURCES.txt +2 -0
  23. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/README.md +0 -0
  24. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/pyproject.toml +0 -0
  25. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/ingest/__init__.py +0 -0
  26. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/models/__init__.py +0 -0
  27. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/models/resnet_vae.py +0 -0
  28. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/models/resnetlight_decoder.py +0 -0
  29. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/models/resnetlight_encoder.py +0 -0
  30. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/models/simple_vae.py +0 -0
  31. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/models/vae.py +0 -0
  32. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/train/__init__.py +0 -0
  33. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/train/logging.py +0 -0
  34. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/vis/__init__.py +0 -0
  35. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/vis/features.py +0 -0
  36. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima/vis/patchexamples.py +0 -0
  37. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima_spatial.egg-info/dependency_links.txt +0 -0
  38. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima_spatial.egg-info/requires.txt +0 -0
  39. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/src/vima_spatial.egg-info/top_level.txt +0 -0
  40. {vima_spatial-0.2.3 → vima_spatial-0.2.4}/tests/test_ra_regression.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: vima-spatial
3
- Version: 0.2.3
3
+ Version: 0.2.4
4
4
  Summary: variational inference-based microniche analysis
5
5
  Home-page: https://github.com/yakirr/vima
6
6
  Author: Yakir Reshef
@@ -1,6 +1,6 @@
1
1
  [metadata]
2
2
  name = vima-spatial
3
- version = 0.2.3
3
+ version = 0.2.4
4
4
  author = Yakir Reshef
5
5
  author_email = yreshef@broadinstitute.org
6
6
  description = variational inference-based microniche analysis
@@ -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__ = ['d', 'cc', 't', 'pp', 'v', 'models',
16
- 'PatchCollection', 'read_samples', 'reindex_by_sid',
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
- pb = lambda x: tqdm(x, ncols=100)
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 pb(eval_loader):
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
- print('applying models')
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
- 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)]
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
- print(f'\033[32mP = {p}\033[0m')
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=False):
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
- print('computing MAT') #TODO: rename MAM to MAT in code if we keep this nomenclature
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 tqdm(ds.modelspecific_fingerprints(), total=ds.nmodels, ncols=100):
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=False, allow_low_sample_size=False,
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][ds.obs.index[kept]]
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
- print('performing association test')
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)
@@ -1,2 +1,3 @@
1
1
  from .patchcollection import *
2
2
  from .samples import *
3
+ from .download import *
@@ -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 tqdm import tqdm
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, verbose=False):
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 pb(samples.values()):
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, verbose=False):
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 verbose:
142
+ if logger.isEnabledFor(logging.DEBUG):
143
143
  fmt = lambda a: ' '.join(f'{v:.2g}' for v in a)
144
- print(f'per-channel means: {fmt(self.means)}')
145
- print(f'per-channel stds: {fmt(self.stds)}')
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, verbose=False):
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
- if verbose: print(f'Normalizing color channels (normalization={normalization})...')
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, verbose=False,
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, verbose=verbose)
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, verbose=verbose)
183
- self.normalize(normalization=normalization, verbose=verbose)
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
- print('WARNING: Data augmentation only available in pytorch mode. Will leave augmentation off')
225
+ logger.warning('Data augmentation only available in pytorch mode. Will leave augmentation off')
228
226
  return
229
- print('\033[90m[PatchCollection: data augmentation on]\033[0m')
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
- print('\033[90m[PatchCollection: data augmentation is off]\033[0m')
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
- print('\033[90m[PatchCollection: in pytorch mode]\033[0m')
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
- print('\033[90m[PatchCollection: in numpy mode]\033[0m')
252
+ logger.debug('[PatchCollection: in numpy mode]')
255
253
 
256
- def subset(self, ix, percentile_thresh, normalization, verbose):
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, verbose=verbose)
264
- self.normalize(normalization=normalization, verbose=verbose)
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 tqdm import tqdm
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 pb(files[:stop_after]):
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 copy
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
- """Return model `i`'s embedding and neighbor graph as an AnnData."""
83
- d = ad.AnnData(X=self._adata.obsm[f'X_{i}'], obs=self._adata.obs.copy())
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 = copy.deepcopy(self._adata.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
- print('Computing UMAP...')
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 pb(range(self.nmodels)):
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
- print(VT.shape)
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