codeconv 1.0.0__py3-none-any.whl
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.
codeconv.py
ADDED
|
@@ -0,0 +1,3277 @@
|
|
|
1
|
+
"""
|
|
2
|
+
CoexpressDeconvolve main module.
|
|
3
|
+
|
|
4
|
+
Multi-slice Visium deconvolution to in-silico single cells.
|
|
5
|
+
Pipeline: load -> density -> ODG -> manifold -> K-sweep -> LDA -> sampling -> placement -> export.
|
|
6
|
+
|
|
7
|
+
Multi-slice handling:
|
|
8
|
+
visium_path can be a string (single slice) or a dict {name: path} / list (multi).
|
|
9
|
+
Per-slice params (min_umi, anchor_mean_factor, low_slice_quality) accept either a
|
|
10
|
+
scalar (broadcast to all slices) or a dict keyed by slice name.
|
|
11
|
+
|
|
12
|
+
LDA is fit per-slice and topics are aligned across slices by Hungarian matching on
|
|
13
|
+
cosine similarity of betas; the consensus beta is the mean of aligned betas.
|
|
14
|
+
Per-slice theta is then refit against the frozen consensus beta via a variational E-step.
|
|
15
|
+
The gene-coexpression manifold is joint over the intersected
|
|
16
|
+
gene set; per-slice beta projection extends consensus topics to each slice's full
|
|
17
|
+
cleaned gene list.
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
import os
|
|
21
|
+
import re
|
|
22
|
+
import glob
|
|
23
|
+
import gzip
|
|
24
|
+
import json
|
|
25
|
+
import shutil
|
|
26
|
+
import time
|
|
27
|
+
import warnings
|
|
28
|
+
from dataclasses import dataclass, field
|
|
29
|
+
from typing import Optional, Union, List, Dict
|
|
30
|
+
|
|
31
|
+
try:
|
|
32
|
+
import numpy as np
|
|
33
|
+
import pandas as pd
|
|
34
|
+
import scipy.io
|
|
35
|
+
import scipy.sparse as sp
|
|
36
|
+
import matplotlib.pyplot as plt
|
|
37
|
+
import seaborn as sns
|
|
38
|
+
# tqdm.auto, not tqdm.notebook: notebook mode raises ImportError (IProgress
|
|
39
|
+
# not found) the moment the module is used outside Jupyter: in a plain
|
|
40
|
+
# script, a CI run, or a terminal session. auto picks the widget bar inside
|
|
41
|
+
# a notebook and the text bar everywhere else.
|
|
42
|
+
from tqdm.auto import tqdm
|
|
43
|
+
from sklearn.decomposition import FastICA, LatentDirichletAllocation
|
|
44
|
+
from sklearn.preprocessing import StandardScaler, normalize
|
|
45
|
+
from scipy.spatial.distance import cdist
|
|
46
|
+
from scipy.spatial import cKDTree
|
|
47
|
+
from scipy.optimize import linear_sum_assignment
|
|
48
|
+
from scipy.special import digamma, polygamma
|
|
49
|
+
import h5py
|
|
50
|
+
import umap
|
|
51
|
+
except ImportError as e:
|
|
52
|
+
missing_pkg = str(e).split()[-1]
|
|
53
|
+
raise ImportError(
|
|
54
|
+
f"Missing package: {missing_pkg}. "
|
|
55
|
+
f"Please install required dependencies via:\n"
|
|
56
|
+
f"pip install numpy pandas scipy matplotlib seaborn scikit-learn h5py tqdm umap-learn"
|
|
57
|
+
)
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
# --- Profiling & figure-output helpers (added for STAR Protocol reproducibility) ---
|
|
61
|
+
import sys as _sys
|
|
62
|
+
import threading as _threading
|
|
63
|
+
import functools as _functools
|
|
64
|
+
try:
|
|
65
|
+
import resource as _resource
|
|
66
|
+
except Exception: # pragma: no cover
|
|
67
|
+
_resource = None
|
|
68
|
+
|
|
69
|
+
# Every step figure is written here (600 dpi) in addition to inline display.
|
|
70
|
+
# Override via set_figure_output().
|
|
71
|
+
FIGURE_DIR = "./figures"
|
|
72
|
+
FIGURE_DPI = 600
|
|
73
|
+
|
|
74
|
+
# Set by step5_ksweep to the number of topics it recommends (or to the manual
|
|
75
|
+
# n_topics= override, if one was given), so Step 6 reads:
|
|
76
|
+
# model = codeconv.step6_final_deconvolution(..., n_topics=codeconv.recommended_K)
|
|
77
|
+
# None until a sweep has been run.
|
|
78
|
+
recommended_K: Optional[int] = None
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def set_figure_output(directory: str = "./figures", dpi: int = 600):
|
|
82
|
+
"""Configure where step figures are written and at what resolution."""
|
|
83
|
+
global FIGURE_DIR, FIGURE_DPI
|
|
84
|
+
FIGURE_DIR = str(directory)
|
|
85
|
+
FIGURE_DPI = int(dpi)
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
def _show(fig_name: str):
|
|
89
|
+
"""Save the current figure to FIGURE_DIR at FIGURE_DPI, then show it.
|
|
90
|
+
|
|
91
|
+
Drop-in replacement for plt.show(): inline notebook display is unchanged, but a
|
|
92
|
+
high-resolution PNG is also written to ./figures for the manuscript.
|
|
93
|
+
"""
|
|
94
|
+
try:
|
|
95
|
+
os.makedirs(FIGURE_DIR, exist_ok=True)
|
|
96
|
+
safe = re.sub(r"[^A-Za-z0-9._-]+", "_", str(fig_name)).strip("_")
|
|
97
|
+
plt.savefig(os.path.join(FIGURE_DIR, f"{safe}.png"),
|
|
98
|
+
dpi=FIGURE_DPI, bbox_inches="tight")
|
|
99
|
+
except Exception as exc: # never let figure I/O break a pipeline run
|
|
100
|
+
print(f" [figure not saved: {exc}]")
|
|
101
|
+
plt.show()
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
def _current_rss_bytes() -> int:
|
|
105
|
+
"""Best-effort current resident-set size of this process, in bytes."""
|
|
106
|
+
try:
|
|
107
|
+
import psutil
|
|
108
|
+
return int(psutil.Process().memory_info().rss)
|
|
109
|
+
except Exception:
|
|
110
|
+
pass
|
|
111
|
+
try:
|
|
112
|
+
with open("/proc/self/statm") as fh: # Linux
|
|
113
|
+
pages = int(fh.read().split()[1])
|
|
114
|
+
return pages * os.sysconf("SC_PAGE_SIZE")
|
|
115
|
+
except Exception:
|
|
116
|
+
pass
|
|
117
|
+
if _resource is not None:
|
|
118
|
+
maxrss = _resource.getrusage(_resource.RUSAGE_SELF).ru_maxrss
|
|
119
|
+
return int(maxrss) if _sys.platform == "darwin" else int(maxrss) * 1024
|
|
120
|
+
return 0
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
class _PeakRSS:
|
|
124
|
+
"""Sample process RSS in a background thread to capture a per-step peak."""
|
|
125
|
+
|
|
126
|
+
def __init__(self, interval: float = 0.1):
|
|
127
|
+
self.interval = interval
|
|
128
|
+
self.peak = 0
|
|
129
|
+
self._stop = _threading.Event()
|
|
130
|
+
self._thread = None
|
|
131
|
+
|
|
132
|
+
def start(self):
|
|
133
|
+
self.peak = _current_rss_bytes()
|
|
134
|
+
self._thread = _threading.Thread(target=self._loop, daemon=True)
|
|
135
|
+
self._thread.start()
|
|
136
|
+
return self
|
|
137
|
+
|
|
138
|
+
def _loop(self):
|
|
139
|
+
while not self._stop.wait(self.interval):
|
|
140
|
+
rss = _current_rss_bytes()
|
|
141
|
+
if rss > self.peak:
|
|
142
|
+
self.peak = rss
|
|
143
|
+
|
|
144
|
+
def stop(self) -> int:
|
|
145
|
+
self._stop.set()
|
|
146
|
+
if self._thread is not None:
|
|
147
|
+
self._thread.join(timeout=1.0)
|
|
148
|
+
rss = _current_rss_bytes()
|
|
149
|
+
if rss > self.peak:
|
|
150
|
+
self.peak = rss
|
|
151
|
+
return self.peak
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def _profile_step(fn):
|
|
155
|
+
"""Decorator: report peak process RAM for a step, mirroring the per-step timing."""
|
|
156
|
+
@_functools.wraps(fn)
|
|
157
|
+
def wrapper(*args, **kwargs):
|
|
158
|
+
monitor = _PeakRSS().start()
|
|
159
|
+
try:
|
|
160
|
+
return fn(*args, **kwargs)
|
|
161
|
+
finally:
|
|
162
|
+
peak_gb = monitor.stop() / (1024 ** 3)
|
|
163
|
+
print(f" [{fn.__name__}] peak RAM: {peak_gb:.2f} GB")
|
|
164
|
+
return wrapper
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
# Module-level RNG state. Set via codeconv.set_seed(int) at the top of the notebook.
|
|
168
|
+
_SEED = 42
|
|
169
|
+
_RNG = np.random.default_rng(_SEED)
|
|
170
|
+
|
|
171
|
+
|
|
172
|
+
def set_seed(seed: int):
|
|
173
|
+
"""Set the module-wide random seed. Call once at the top of the notebook."""
|
|
174
|
+
global _SEED, _RNG
|
|
175
|
+
_SEED = int(seed)
|
|
176
|
+
_RNG = np.random.default_rng(_SEED)
|
|
177
|
+
np.random.seed(_SEED)
|
|
178
|
+
print(f"codeconv: random seed set to {_SEED}")
|
|
179
|
+
|
|
180
|
+
|
|
181
|
+
# Container types
|
|
182
|
+
|
|
183
|
+
@dataclass
|
|
184
|
+
class SliceData:
|
|
185
|
+
"""Per-slice data accumulator. Steps progressively fill in fields."""
|
|
186
|
+
name: str
|
|
187
|
+
spatial_path: str
|
|
188
|
+
counts: sp.csr_matrix
|
|
189
|
+
gene_names: List[str]
|
|
190
|
+
barcodes: List[str]
|
|
191
|
+
coords: np.ndarray
|
|
192
|
+
total_umi: np.ndarray
|
|
193
|
+
scale_factors: dict
|
|
194
|
+
# source platform key, see PLATFORM_PROFILES
|
|
195
|
+
platform: str = 'visium'
|
|
196
|
+
# filled by Step 2
|
|
197
|
+
n_cells: Optional[np.ndarray] = None
|
|
198
|
+
engine_params: Optional[dict] = None
|
|
199
|
+
# filled by Step 3 per-slice (post noise filter, slice-specific)
|
|
200
|
+
counts_clean: Optional[sp.csr_matrix] = None
|
|
201
|
+
genes_clean: Optional[List[str]] = None
|
|
202
|
+
# filled by Step 6 per-slice
|
|
203
|
+
theta: Optional[np.ndarray] = None
|
|
204
|
+
beta_final: Optional[np.ndarray] = None
|
|
205
|
+
|
|
206
|
+
|
|
207
|
+
@dataclass
|
|
208
|
+
class OdgPack:
|
|
209
|
+
"""Output of Step 3. Joint overdispersed-gene selection on the intersected gene set."""
|
|
210
|
+
intersected_genes: List[str]
|
|
211
|
+
odg_names: List[str]
|
|
212
|
+
odg_per_slice: Dict[str, sp.csr_matrix]
|
|
213
|
+
odg_concat: sp.csr_matrix
|
|
214
|
+
species: str
|
|
215
|
+
|
|
216
|
+
|
|
217
|
+
@dataclass
|
|
218
|
+
class Manifold:
|
|
219
|
+
"""Output of Step 4. Joint manifold on intersected genes."""
|
|
220
|
+
embedding: np.ndarray
|
|
221
|
+
intersected_genes: List[str]
|
|
222
|
+
odg_indices_in_intersected: List[int]
|
|
223
|
+
species: str
|
|
224
|
+
|
|
225
|
+
|
|
226
|
+
@dataclass
|
|
227
|
+
class Model:
|
|
228
|
+
"""Output of Step 6. Consensus beta + per-slice theta."""
|
|
229
|
+
n_topics: int
|
|
230
|
+
odg_names: List[str]
|
|
231
|
+
beta_consensus: np.ndarray
|
|
232
|
+
qc_df: pd.DataFrame
|
|
233
|
+
per_slice_betas: Dict[str, np.ndarray]
|
|
234
|
+
per_slice_stability: Optional[Dict[str, float]] = None
|
|
235
|
+
|
|
236
|
+
|
|
237
|
+
@dataclass
|
|
238
|
+
class KSweepResult:
|
|
239
|
+
"""Output of Step 5. Compact repr so notebook auto-display stays one line.
|
|
240
|
+
|
|
241
|
+
All metrics remain accessible as attributes for programmatic inspection,
|
|
242
|
+
custom plotting, or persistence; only the default display is suppressed.
|
|
243
|
+
"""
|
|
244
|
+
perplexity: Dict[str, Dict[int, float]]
|
|
245
|
+
rare_topics: Dict[str, Dict[int, int]]
|
|
246
|
+
alpha_mean: Dict[str, Dict[int, float]]
|
|
247
|
+
alpha_per_topic: Dict[str, Dict[int, np.ndarray]]
|
|
248
|
+
perc_rare_thresh: float
|
|
249
|
+
recommended_k: Dict[str, Dict[str, object]]
|
|
250
|
+
k_tol: float = 0.01
|
|
251
|
+
k_band: float = 0.05
|
|
252
|
+
k_rule: str = "percentile"
|
|
253
|
+
chosen_k: Optional[int] = None
|
|
254
|
+
|
|
255
|
+
def __repr__(self) -> str:
|
|
256
|
+
if not self.perplexity:
|
|
257
|
+
return "KSweepResult(empty)"
|
|
258
|
+
targets = list(self.perplexity.keys())
|
|
259
|
+
any_label = targets[0]
|
|
260
|
+
ks = sorted(self.perplexity[any_label].keys())
|
|
261
|
+
if ks:
|
|
262
|
+
k_span = f"K={ks[0]}..{ks[-1]}"
|
|
263
|
+
else:
|
|
264
|
+
k_span = "K=[]"
|
|
265
|
+
rec_txt = ""
|
|
266
|
+
if self.recommended_k:
|
|
267
|
+
label = '_joint' if '_joint' in self.recommended_k else next(iter(self.recommended_k))
|
|
268
|
+
rec_txt = f", recommended_K={self.recommended_k[label]['k']}"
|
|
269
|
+
chosen_txt = f", chosen_K={self.chosen_k}" if self.chosen_k is not None else ""
|
|
270
|
+
return f"KSweepResult(targets={targets}, {k_span}{rec_txt}{chosen_txt})"
|
|
271
|
+
|
|
272
|
+
@property
|
|
273
|
+
def k(self) -> Optional[int]:
|
|
274
|
+
"""The single recommended K to carry into Step 6.
|
|
275
|
+
|
|
276
|
+
Uses the joint sweep when several slices were swept, otherwise the only
|
|
277
|
+
slice. ``chosen_k`` (a manual override passed to step5_ksweep) always
|
|
278
|
+
wins, so ``ksweep.k`` is the value the run actually proceeds with.
|
|
279
|
+
"""
|
|
280
|
+
if self.chosen_k is not None:
|
|
281
|
+
return int(self.chosen_k)
|
|
282
|
+
if not self.recommended_k:
|
|
283
|
+
return None
|
|
284
|
+
label = '_joint' if '_joint' in self.recommended_k else next(iter(self.recommended_k))
|
|
285
|
+
return int(self.recommended_k[label]['k'])
|
|
286
|
+
|
|
287
|
+
def summary(self) -> pd.DataFrame:
|
|
288
|
+
"""Per-K sweep metrics as a tidy DataFrame, for the record or for replotting."""
|
|
289
|
+
rows = []
|
|
290
|
+
for label, perps in self.perplexity.items():
|
|
291
|
+
p = np.array([perps[k] for k in sorted(perps)], dtype=float)
|
|
292
|
+
span = float(p.max() - p.min())
|
|
293
|
+
in_set = set(self.recommended_k.get(label, {}).get('k_in_band', []))
|
|
294
|
+
for i, k in enumerate(sorted(perps)):
|
|
295
|
+
rows.append({
|
|
296
|
+
'target': label,
|
|
297
|
+
'K': k,
|
|
298
|
+
'perplexity': float(p[i]),
|
|
299
|
+
'perplexity_norm': float((p[i] - p.min()) / span) if span > 0 else 0.0,
|
|
300
|
+
'pct_above_best': float(p[i] / p.min() - 1.0) if p.min() > 0 else 0.0,
|
|
301
|
+
'in_low_band': bool(k in in_set),
|
|
302
|
+
'rare_topics': int(self.rare_topics[label][k]),
|
|
303
|
+
'alpha_mean': float(self.alpha_mean[label][k]),
|
|
304
|
+
})
|
|
305
|
+
return pd.DataFrame(rows)
|
|
306
|
+
|
|
307
|
+
|
|
308
|
+
def recommend_k(
|
|
309
|
+
perplexity: Dict[int, float],
|
|
310
|
+
rule: str = "percentile",
|
|
311
|
+
tol: float = 0.01,
|
|
312
|
+
band: float = 0.05,
|
|
313
|
+
rare_topics: Optional[Dict[int, int]] = None,
|
|
314
|
+
alpha_mean: Optional[Dict[int, float]] = None,
|
|
315
|
+
) -> Dict[str, object]:
|
|
316
|
+
"""Pick the number of topics K from a held-out-perplexity sweep, automatically.
|
|
317
|
+
|
|
318
|
+
The protocol's guidance is "take the highest K that still sits in the
|
|
319
|
+
low-perplexity regime, because rare topics are wanted here": rare topics
|
|
320
|
+
correspond to minor cell populations we intend to recover as individual
|
|
321
|
+
cells, not to over-splitting. This function is that sentence written as a
|
|
322
|
+
rule, so the choice stops being a visual judgement call.
|
|
323
|
+
|
|
324
|
+
Every rule returns the **largest** K inside a low-perplexity set; they differ
|
|
325
|
+
only in how that set is defined.
|
|
326
|
+
|
|
327
|
+
Rule ``percentile`` (default, ``band = 0.05``). Keep the K values whose
|
|
328
|
+
held-out perplexity falls in the best ``band`` fraction of the swept
|
|
329
|
+
perplexities by rank, the 95th-percentile rule, and return the largest.
|
|
330
|
+
At band = 0.05 the cut sits essentially at the minimum, so in practice this
|
|
331
|
+
reads as "among the K values statistically tied with the best fit, take the
|
|
332
|
+
most resolved one". Widening the band past ~0.10 starts admitting K values
|
|
333
|
+
that are visibly worse, which is what over-calls K.
|
|
334
|
+
|
|
335
|
+
Rule ``relative_tolerance``. Keep every K whose held-out perplexity is within
|
|
336
|
+
``tol`` of the best K,
|
|
337
|
+
|
|
338
|
+
P(K) <= min(P) * (1 + tol) tol = 0.01 -> within 1%
|
|
339
|
+
|
|
340
|
+
and return the largest. Its advantage is that it is invariant to the sweep's
|
|
341
|
+
extent by construction: the answer does not move when ``min_k`` or ``max_k``
|
|
342
|
+
change, only when the fit genuinely changes, and it is directly
|
|
343
|
+
interpretable in units of fit quality rather than rank. Worth comparing
|
|
344
|
+
against the default on any new dataset.
|
|
345
|
+
|
|
346
|
+
Rule ``largest_in_band``. Min-max normalize the sweep,
|
|
347
|
+
``p_norm(K) = (P(K) - min P) / (max P - min P)``, and keep every K with
|
|
348
|
+
``p_norm(K) <= band``. Intuitive to read off the plot, but the normalization
|
|
349
|
+
is anchored to ``max P``, which is whatever the smallest swept K happened to
|
|
350
|
+
score, so the recommendation shifts when the sweep range changes even though
|
|
351
|
+
the curve has not. Kept for reading the shaded band on Figure 5.
|
|
352
|
+
|
|
353
|
+
Rule ``lowest_perplexity``. Plain argmin, the conventional choice, kept for
|
|
354
|
+
comparison. It cannot return a K larger than the best-fitting one, so it
|
|
355
|
+
systematically under-calls K in this setting.
|
|
356
|
+
|
|
357
|
+
Rule ``conservative``. Lowest perplexity among K with mean Dirichlet
|
|
358
|
+
alpha < 1 and no rare topics; needs ``rare_topics`` and ``alpha_mean``. This
|
|
359
|
+
was the module's original rule. It optimizes for topics that are all
|
|
360
|
+
well-populated, which is the right target for proportion-style deconvolution
|
|
361
|
+
and the wrong one here, so it is no longer the default.
|
|
362
|
+
|
|
363
|
+
Whatever the rule, a recommendation equal to ``max_k`` means perplexity had
|
|
364
|
+
not started to rise again and the sweep is too narrow to bound K from above;
|
|
365
|
+
the returned ``at_max_k`` flag says so.
|
|
366
|
+
|
|
367
|
+
Parameters
|
|
368
|
+
----------
|
|
369
|
+
perplexity : Dict[int, float]
|
|
370
|
+
Held-out perplexity keyed by K, e.g. ``ksweep.perplexity['Glioblastoma']``.
|
|
371
|
+
rule : str
|
|
372
|
+
``percentile`` (default) | ``relative_tolerance`` | ``largest_in_band`` |
|
|
373
|
+
``lowest_perplexity`` | ``conservative``.
|
|
374
|
+
tol : float
|
|
375
|
+
Relative tolerance on the best perplexity, for ``relative_tolerance``.
|
|
376
|
+
Default 0.01 (1%). Lower it for a stricter, smaller K.
|
|
377
|
+
band : float
|
|
378
|
+
Band width as a fraction, for ``percentile`` and ``largest_in_band``.
|
|
379
|
+
Default 0.05 (the 95th-percentile rule). Raise it for a larger K.
|
|
380
|
+
rare_topics, alpha_mean : Optional[Dict[int, float]]
|
|
381
|
+
Required by ``conservative`` only.
|
|
382
|
+
|
|
383
|
+
Returns
|
|
384
|
+
-------
|
|
385
|
+
dict
|
|
386
|
+
``k``, ``criterion``, ``rule``, ``tol``, ``band``, ``at_max_k``,
|
|
387
|
+
``k_in_band`` (every K in the low-perplexity set) and ``perplexity_norm``.
|
|
388
|
+
|
|
389
|
+
Examples
|
|
390
|
+
--------
|
|
391
|
+
>>> ksweep = codeconv.step5_ksweep(odg_pack, min_k=3, max_k=15)
|
|
392
|
+
>>> recommend_k(ksweep.perplexity['Glioblastoma'], tol=0.01)['k']
|
|
393
|
+
"""
|
|
394
|
+
if not perplexity:
|
|
395
|
+
raise ValueError("recommend_k: empty perplexity sweep.")
|
|
396
|
+
|
|
397
|
+
ks = sorted(perplexity.keys())
|
|
398
|
+
p = np.array([float(perplexity[k]) for k in ks])
|
|
399
|
+
span = float(p.max() - p.min())
|
|
400
|
+
p_norm = (p - p.min()) / span if span > 0 else np.zeros_like(p)
|
|
401
|
+
norm_map = {int(k): float(v) for k, v in zip(ks, p_norm)}
|
|
402
|
+
|
|
403
|
+
if rule == "relative_tolerance":
|
|
404
|
+
if tol < 0.0:
|
|
405
|
+
raise ValueError(f"recommend_k: tol must be >= 0, got {tol}.")
|
|
406
|
+
threshold = float(p.min()) * (1.0 + tol)
|
|
407
|
+
in_band = [int(k) for k, v in zip(ks, p) if v <= threshold]
|
|
408
|
+
k_star = max(in_band)
|
|
409
|
+
criterion = (
|
|
410
|
+
f"largest K with held-out perplexity within {tol:.1%} of the best "
|
|
411
|
+
f"(<= {threshold:.1f}; best = {p.min():.1f} at K = {ks[int(np.argmin(p))]})"
|
|
412
|
+
)
|
|
413
|
+
elif rule == "largest_in_band":
|
|
414
|
+
if not (0.0 <= band <= 1.0):
|
|
415
|
+
raise ValueError(f"recommend_k: band must be in [0, 1], got {band}.")
|
|
416
|
+
in_band = [int(k) for k, v in zip(ks, p_norm) if v <= band]
|
|
417
|
+
if not in_band: # only possible if band < 0 rounding; keep it safe
|
|
418
|
+
in_band = [int(ks[int(np.argmin(p))])]
|
|
419
|
+
k_star = max(in_band)
|
|
420
|
+
criterion = (
|
|
421
|
+
f"largest K with normalized held-out perplexity <= {band:.2f} "
|
|
422
|
+
f"(low-perplexity band = best {band * 100:.0f}% of the sweep range)"
|
|
423
|
+
)
|
|
424
|
+
elif rule == "percentile":
|
|
425
|
+
if not (0.0 <= band <= 1.0):
|
|
426
|
+
raise ValueError(f"recommend_k: band must be in [0, 1], got {band}.")
|
|
427
|
+
threshold = float(np.percentile(p, band * 100.0))
|
|
428
|
+
in_band = [int(k) for k, v in zip(ks, p) if v <= threshold]
|
|
429
|
+
if not in_band:
|
|
430
|
+
in_band = [int(ks[int(np.argmin(p))])]
|
|
431
|
+
k_star = max(in_band)
|
|
432
|
+
criterion = (
|
|
433
|
+
f"largest K in the best {band:.0%} of swept perplexities by rank "
|
|
434
|
+
f"(<= {threshold:.1f})"
|
|
435
|
+
)
|
|
436
|
+
elif rule == "lowest_perplexity":
|
|
437
|
+
in_band = [int(ks[int(np.argmin(p))])]
|
|
438
|
+
k_star = in_band[0]
|
|
439
|
+
criterion = "lowest held-out perplexity"
|
|
440
|
+
elif rule == "conservative":
|
|
441
|
+
if rare_topics is None or alpha_mean is None:
|
|
442
|
+
raise ValueError("recommend_k: rule='conservative' needs rare_topics and alpha_mean.")
|
|
443
|
+
rare_arr = np.array([rare_topics[k] for k in ks])
|
|
444
|
+
alpha_arr = np.array([alpha_mean[k] for k in ks])
|
|
445
|
+
both = (alpha_arr < 1.0) & (rare_arr == 0)
|
|
446
|
+
if both.any():
|
|
447
|
+
idx = int(np.argmin(np.where(both, p, np.inf)))
|
|
448
|
+
criterion = "alpha<1 AND rare==0; lowest perplexity"
|
|
449
|
+
elif (alpha_arr < 1.0).any():
|
|
450
|
+
idx = int(np.argmin(np.where(alpha_arr < 1.0, p, np.inf)))
|
|
451
|
+
criterion = "alpha<1 only; lowest perplexity (no K reached rare==0)"
|
|
452
|
+
else:
|
|
453
|
+
idx = int(np.argmin(p))
|
|
454
|
+
criterion = "fallback: lowest perplexity (no K reached alpha<1)"
|
|
455
|
+
k_star = int(ks[idx])
|
|
456
|
+
in_band = [k_star]
|
|
457
|
+
else:
|
|
458
|
+
raise ValueError(
|
|
459
|
+
f"recommend_k: unknown rule {rule!r}; expected 'relative_tolerance', "
|
|
460
|
+
f"'largest_in_band', 'percentile', 'lowest_perplexity' or 'conservative'."
|
|
461
|
+
)
|
|
462
|
+
|
|
463
|
+
return {
|
|
464
|
+
'k': int(k_star),
|
|
465
|
+
'criterion': criterion,
|
|
466
|
+
'rule': rule,
|
|
467
|
+
'tol': float(tol),
|
|
468
|
+
'band': float(band),
|
|
469
|
+
'k_in_band': in_band,
|
|
470
|
+
'at_max_k': bool(k_star == ks[-1]),
|
|
471
|
+
'perplexity_norm': norm_map,
|
|
472
|
+
}
|
|
473
|
+
|
|
474
|
+
|
|
475
|
+
def cellchat_spatial_factors(
|
|
476
|
+
cells: Dict[str, dict],
|
|
477
|
+
slices: Dict[str, "SliceData"],
|
|
478
|
+
spot_pitch_um: float = 100.0,
|
|
479
|
+
interaction_range_um: Optional[float] = 50.0,
|
|
480
|
+
) -> Dict[str, Dict[str, float]]:
|
|
481
|
+
"""Derive CellChat's ``spatial.factors`` (ratio, tol) for the reconstructed cells.
|
|
482
|
+
|
|
483
|
+
**The pixel-to-micron ratio is calibrated from the spot pitch, not from
|
|
484
|
+
``spot_diameter_fullres``.** CellChat's vignette computes it as
|
|
485
|
+
``spot.size / spot_diameter_fullres`` with ``spot.size = 65``, described as
|
|
486
|
+
"the theoretical spot size (um) in 10X Visium", while 10x's physical capture
|
|
487
|
+
spot is 55 um. Neither constant reproduces the real geometry: measured on the
|
|
488
|
+
worked example (11 mm CytAssist, ``spot_diameter_fullres = 219.17 px``, true
|
|
489
|
+
spot pitch 362.93 px), 55 um implies a 91 um pitch and 65 um implies 108 um,
|
|
490
|
+
against a documented pitch of 100 um. The pitch is the unambiguous number:
|
|
491
|
+
it is fixed by the array, verifiable from the spot coordinates, and identical
|
|
492
|
+
across Visium capture areas, so the ratio is derived from it directly::
|
|
493
|
+
|
|
494
|
+
ratio = spot_pitch_um / median nearest-neighbour distance between spot centres
|
|
495
|
+
|
|
496
|
+
which for the worked example gives 0.2755 um/pixel and implies
|
|
497
|
+
``spot_diameter_fullres`` spans ~60 um, i.e. between the two conventions.
|
|
498
|
+
|
|
499
|
+
``tol``, half the characteristic centre-to-centre distance. CellChat suggests
|
|
500
|
+
``spot_size/2`` for Visium spots, but after deconvolution the objects are
|
|
501
|
+
reconstructed cells packed *inside* each spot, so the characteristic spacing
|
|
502
|
+
is the cell-to-cell nearest-neighbour distance, not the spot pitch. Returned
|
|
503
|
+
as half the median nearest-neighbour distance, following the documented
|
|
504
|
+
fallback ("if the cell/spot size is not known ... tol can be the half value of
|
|
505
|
+
the minimum centre-to-centre distance"). Passing the spot-scale
|
|
506
|
+
``55/2 = 27.5 um`` instead silently widens every range by that amount.
|
|
507
|
+
|
|
508
|
+
Sub-spot geometry check. Within a spot, cells are positioned by the
|
|
509
|
+
Fibonacci/Vogel packing of ``step8_geometry_and_placement``, not measured. The
|
|
510
|
+
question that matters is therefore not how many neighbours share a spot, but
|
|
511
|
+
whether the packing **decides which pairs count**. It does not, as long as the
|
|
512
|
+
cutoff covers the whole spot footprint: then every intra-spot pair is inside
|
|
513
|
+
the range whatever the packing did with it, and the inference reduces to
|
|
514
|
+
co-occupancy in the same capture spot, which is measured. So the diagnostic
|
|
515
|
+
reported here is ``intra_spot_pairs_captured``, the fraction of within-spot
|
|
516
|
+
cell pairs falling inside ``interaction_range_um + tol``. At 100% the result is
|
|
517
|
+
invariant to the packing; below ~95% the packing starts selecting pairs and
|
|
518
|
+
short-range results become an artefact of the placement model. On the worked
|
|
519
|
+
example (spot footprint 60.4 um) the coverage is 100% from
|
|
520
|
+
``interaction_range_um = 50`` upward and falls to 81% at 30 um.
|
|
521
|
+
|
|
522
|
+
``same_spot_fraction`` is also returned, as context for how much of the signal
|
|
523
|
+
is spot co-occupancy versus between-spot architecture, but it is not an error
|
|
524
|
+
condition.
|
|
525
|
+
|
|
526
|
+
What this function deliberately does **not** return is ``scale.distance``.
|
|
527
|
+
That factor is applied to ``d.spatial``, which in ``computeCommunProb`` is a
|
|
528
|
+
*cell-group x cell-group* matrix of trimmed-mean neighbour distances built by
|
|
529
|
+
``computeRegionDistance``, not a cell-to-cell distance matrix. Its minimum
|
|
530
|
+
therefore depends on the cell-type labels, which do not exist until the
|
|
531
|
+
annotation step in R. Derive it there, with CellChat's own function::
|
|
532
|
+
|
|
533
|
+
res <- computeRegionDistance(coordinates = spatial.locs, meta = meta.t,
|
|
534
|
+
interaction.range = 100, ratio = ratio, tol = tol,
|
|
535
|
+
contact.dependent = TRUE, contact.range = 100)
|
|
536
|
+
scale.distance <- 1.5 / min(res$d.spatial, na.rm = TRUE)
|
|
537
|
+
|
|
538
|
+
Call after ``step8_geometry_and_placement``. ``step9_export_results`` also
|
|
539
|
+
calls it and records the values in ``run_summary.json``.
|
|
540
|
+
|
|
541
|
+
Parameters
|
|
542
|
+
----------
|
|
543
|
+
spot_pitch_um : float
|
|
544
|
+
Centre-to-centre distance between capture spots, in microns. 100.0 for
|
|
545
|
+
every 10x Visium capture area; pass the array pitch for other platforms.
|
|
546
|
+
interaction_range_um : Optional[float]
|
|
547
|
+
The CellChat ``interaction.range`` you intend to use, in microns.
|
|
548
|
+
Default 50.0: with a 60 um Visium spot footprint this is the shortest
|
|
549
|
+
range that still captures every intra-spot pair, so it is the most
|
|
550
|
+
conservative packing-invariant choice. Diagnostics are reported at a
|
|
551
|
+
cutoff of ``interaction_range_um + tol``. Pass None to skip them.
|
|
552
|
+
|
|
553
|
+
Returns
|
|
554
|
+
-------
|
|
555
|
+
Dict[str, Dict[str, float]]
|
|
556
|
+
Per slice: ``ratio`` (um per pixel), ``tol``, ``min_distance_um``,
|
|
557
|
+
``median_nn_distance_um``, ``spot_pitch_px``, ``spot_diameter_um``,
|
|
558
|
+
``n_cells``, and, unless ``interaction_range_um`` is None,
|
|
559
|
+
``interaction_range_um``, ``cutoff_um``, ``mean_neighbours``,
|
|
560
|
+
``same_spot_fraction`` and ``intra_spot_pairs_captured``.
|
|
561
|
+
"""
|
|
562
|
+
out: Dict[str, Dict[str, float]] = {}
|
|
563
|
+
for name, payload in cells.items():
|
|
564
|
+
sd = slices[name]
|
|
565
|
+
coords = payload.get('final_coords')
|
|
566
|
+
if coords is None:
|
|
567
|
+
raise ValueError(
|
|
568
|
+
f"cellchat_spatial_factors [{name}]: cells have no coordinates yet: "
|
|
569
|
+
f"run step8_geometry_and_placement first."
|
|
570
|
+
)
|
|
571
|
+
if np.asarray(coords).shape[0] < 2:
|
|
572
|
+
raise ValueError(f"cellchat_spatial_factors [{name}]: need at least 2 cells.")
|
|
573
|
+
|
|
574
|
+
# Ratio from the spot pitch, measured on the original spot centres.
|
|
575
|
+
spot_xy = np.asarray(sd.coords[:, :2], dtype=float)
|
|
576
|
+
if spot_xy.shape[0] < 2:
|
|
577
|
+
raise ValueError(f"cellchat_spatial_factors [{name}]: need at least 2 spots.")
|
|
578
|
+
d_spot, _ = cKDTree(spot_xy).query(spot_xy, k=2)
|
|
579
|
+
pitch_px = float(np.median(d_spot[:, 1]))
|
|
580
|
+
if not np.isfinite(pitch_px) or pitch_px <= 0:
|
|
581
|
+
raise ValueError(f"cellchat_spatial_factors [{name}]: degenerate spot pitch.")
|
|
582
|
+
ratio = float(spot_pitch_um) / pitch_px
|
|
583
|
+
|
|
584
|
+
xy_um = np.asarray(coords[:, :2], dtype=float) * ratio
|
|
585
|
+
|
|
586
|
+
# k-d tree nearest-neighbour query, not a dense distance matrix: at
|
|
587
|
+
# 45,530 cells the dense form would be ~16.6 GB.
|
|
588
|
+
tree = cKDTree(xy_um)
|
|
589
|
+
k_query = min(8, xy_um.shape[0])
|
|
590
|
+
dists, _ = tree.query(xy_um, k=k_query)
|
|
591
|
+
neighbours = dists[:, 1:]
|
|
592
|
+
nonzero = neighbours[neighbours > 0]
|
|
593
|
+
if nonzero.size == 0:
|
|
594
|
+
raise ValueError(
|
|
595
|
+
f"cellchat_spatial_factors [{name}]: all cells share coordinates."
|
|
596
|
+
)
|
|
597
|
+
d_min = float(nonzero.min())
|
|
598
|
+
per_cell_nn = np.where(neighbours[:, 0] > 0, neighbours[:, 0], np.nan)
|
|
599
|
+
d_med = float(np.nanmedian(per_cell_nn))
|
|
600
|
+
tol = d_med / 2.0
|
|
601
|
+
|
|
602
|
+
spot_diameter_um = float(sd.scale_factors['spot_diameter_fullres']) * ratio
|
|
603
|
+
|
|
604
|
+
res = {
|
|
605
|
+
'ratio': ratio,
|
|
606
|
+
'tol': tol,
|
|
607
|
+
'min_distance_um': d_min,
|
|
608
|
+
'median_nn_distance_um': d_med,
|
|
609
|
+
'spot_pitch_px': pitch_px,
|
|
610
|
+
'spot_diameter_um': spot_diameter_um,
|
|
611
|
+
'n_cells': int(xy_um.shape[0]),
|
|
612
|
+
}
|
|
613
|
+
print(
|
|
614
|
+
f"CellChat [{name}]: spot pitch {pitch_px:.1f} px = {spot_pitch_um:.0f} um "
|
|
615
|
+
f"-> ratio={ratio:.5f} um/pixel spot footprint {spot_diameter_um:.1f} um "
|
|
616
|
+
f"cell nearest-neighbour distance min={d_min:.1f} median={d_med:.1f} um "
|
|
617
|
+
f"-> spatial.factors(ratio={ratio:.5f}, tol={tol:.2f})"
|
|
618
|
+
)
|
|
619
|
+
|
|
620
|
+
if interaction_range_um is not None:
|
|
621
|
+
cutoff = float(interaction_range_um) + tol
|
|
622
|
+
counts = np.asarray(
|
|
623
|
+
tree.query_ball_point(xy_um, r=cutoff, return_length=True)
|
|
624
|
+
) - 1
|
|
625
|
+
|
|
626
|
+
same_frac = float('nan')
|
|
627
|
+
intra_cov = float('nan')
|
|
628
|
+
barcodes = payload.get('final_barcodes')
|
|
629
|
+
if barcodes is not None:
|
|
630
|
+
parent = np.array([str(b).split('-Topic-')[0] for b in barcodes])
|
|
631
|
+
_, spot_of_cell = np.unique(parent, return_inverse=True)
|
|
632
|
+
|
|
633
|
+
# Does the packing decide which intra-spot pairs count? Spots hold
|
|
634
|
+
# a handful of cells, so the exact per-spot pair enumeration is cheap.
|
|
635
|
+
order = np.argsort(spot_of_cell, kind='stable')
|
|
636
|
+
bounds = np.cumsum(np.bincount(spot_of_cell))[:-1]
|
|
637
|
+
n_pairs = n_inside = 0
|
|
638
|
+
for grp in np.split(order, bounds):
|
|
639
|
+
if grp.size < 2:
|
|
640
|
+
continue
|
|
641
|
+
pts = xy_um[grp]
|
|
642
|
+
dd = np.linalg.norm(pts[:, None, :] - pts[None, :, :], axis=-1)
|
|
643
|
+
iu = np.triu_indices(grp.size, 1)
|
|
644
|
+
n_pairs += iu[0].size
|
|
645
|
+
n_inside += int((dd[iu] <= cutoff).sum())
|
|
646
|
+
intra_cov = n_inside / n_pairs if n_pairs else float('nan')
|
|
647
|
+
|
|
648
|
+
# Same-spot share of neighbours, sampled to bound the cost.
|
|
649
|
+
rng_s = np.random.default_rng(0)
|
|
650
|
+
probe = rng_s.choice(xy_um.shape[0], min(4000, xy_um.shape[0]), replace=False)
|
|
651
|
+
same = tot = 0
|
|
652
|
+
for i in probe:
|
|
653
|
+
nb = [j for j in tree.query_ball_point(xy_um[i], r=cutoff) if j != i]
|
|
654
|
+
tot += len(nb)
|
|
655
|
+
if nb:
|
|
656
|
+
same += int((spot_of_cell[nb] == spot_of_cell[i]).sum())
|
|
657
|
+
same_frac = same / tot if tot else float('nan')
|
|
658
|
+
|
|
659
|
+
res.update({
|
|
660
|
+
'interaction_range_um': float(interaction_range_um),
|
|
661
|
+
'cutoff_um': cutoff,
|
|
662
|
+
'mean_neighbours': float(counts.mean()),
|
|
663
|
+
'same_spot_fraction': same_frac,
|
|
664
|
+
'intra_spot_pairs_captured': intra_cov,
|
|
665
|
+
})
|
|
666
|
+
print(
|
|
667
|
+
f" interaction.range={interaction_range_um:g} um "
|
|
668
|
+
f"(effective cutoff {cutoff:.1f} um): {counts.mean():.1f} neighbours per cell, "
|
|
669
|
+
f"{same_frac:.0%} of them in the same spot, "
|
|
670
|
+
f"{intra_cov:.1%} of intra-spot pairs captured"
|
|
671
|
+
)
|
|
672
|
+
if np.isfinite(intra_cov) and intra_cov < 0.95:
|
|
673
|
+
print(
|
|
674
|
+
f" WARNING: only {intra_cov:.0%} of within-spot cell pairs fall inside this "
|
|
675
|
+
f"range, so the Fibonacci/Vogel packing of Step 8 is selecting which pairs are "
|
|
676
|
+
f"counted and the result is partly an artefact of the placement model. Use "
|
|
677
|
+
f"interaction.range >= {spot_diameter_um - tol:.0f} um to cover the whole "
|
|
678
|
+
f"{spot_diameter_um:.0f} um spot footprint and keep the inference "
|
|
679
|
+
f"packing-invariant."
|
|
680
|
+
)
|
|
681
|
+
out[name] = res
|
|
682
|
+
return out
|
|
683
|
+
|
|
684
|
+
|
|
685
|
+
# --- Platform profiles ---------------------------------------------------------
|
|
686
|
+
#
|
|
687
|
+
# Everything downstream of Step 1 only needs four things per slice: a counts
|
|
688
|
+
# matrix, per-barcode coordinates, a spot footprint, and a spot pitch. Visium
|
|
689
|
+
# supplies all four in the SpaceRanger layout; other spot-based platforms supply
|
|
690
|
+
# the first two and leave the geometry to be filled in from what is known about
|
|
691
|
+
# the array. This table is that knowledge.
|
|
692
|
+
#
|
|
693
|
+
# One calibration rule covers every platform: the physical scale of the
|
|
694
|
+
# coordinates is recovered from the array's *known* centre-to-centre spacing
|
|
695
|
+
# divided by the spacing actually measured in the file,
|
|
696
|
+
#
|
|
697
|
+
# um_per_unit = pitch_um / median nearest-neighbour distance between capture units
|
|
698
|
+
#
|
|
699
|
+
# This avoids depending on any per-vendor field whose meaning is ambiguous. On
|
|
700
|
+
# the worked Visium example it recovers 0.2755 um/pixel, against the array's
|
|
701
|
+
# documented 100 um spot pitch.
|
|
702
|
+
#
|
|
703
|
+
# `coord_units`:
|
|
704
|
+
# 'grid' , coordinates are integer array indices; multiply by `pitch_um`.
|
|
705
|
+
# 'scaled', coordinates are in a linear unit of unknown scale; calibrate with
|
|
706
|
+
# the rule above.
|
|
707
|
+
PLATFORM_PROFILES: dict = {
|
|
708
|
+
'visium': {
|
|
709
|
+
'label': '10x Visium',
|
|
710
|
+
'spot_diameter_um': 55.0,
|
|
711
|
+
'pitch_um': 100.0,
|
|
712
|
+
'coord_units': 'scaled',
|
|
713
|
+
'has_image': True,
|
|
714
|
+
},
|
|
715
|
+
'dbit_seq': {
|
|
716
|
+
'label': 'DBiT-seq',
|
|
717
|
+
# Microchannel width; the 10 / 25 / 50 um variants are distinguished by
|
|
718
|
+
# `pixel_um`, read from the sample name when it carries one. Pitch is
|
|
719
|
+
# twice the width, because channels alternate with equal spacing.
|
|
720
|
+
'spot_diameter_um': 50.0,
|
|
721
|
+
'pitch_um': 100.0,
|
|
722
|
+
'coord_units': 'grid',
|
|
723
|
+
'has_image': False,
|
|
724
|
+
},
|
|
725
|
+
}
|
|
726
|
+
|
|
727
|
+
|
|
728
|
+
def detect_platform(data_path: str) -> str:
|
|
729
|
+
"""Identify which spot-based platform a dataset directory or file came from.
|
|
730
|
+
|
|
731
|
+
Detection is by layout, in order of specificity:
|
|
732
|
+
|
|
733
|
+
================== =========================================================
|
|
734
|
+
``visium`` a ``spatial/`` directory with ``scalefactors_json.json``
|
|
735
|
+
``dbit_seq`` integer ``AxB`` labels on the 50 x 50 microfluidic grid
|
|
736
|
+
================== =========================================================
|
|
737
|
+
|
|
738
|
+
Returns the platform key; raises if nothing matches. Detection is a
|
|
739
|
+
heuristic; pass ``platform=`` to ``step1_acquisition_and_anchoring`` to
|
|
740
|
+
override it.
|
|
741
|
+
"""
|
|
742
|
+
if os.path.isdir(data_path):
|
|
743
|
+
if os.path.exists(os.path.join(data_path, 'spatial', 'scalefactors_json.json')):
|
|
744
|
+
return 'visium'
|
|
745
|
+
tables = _find_count_tables(data_path)
|
|
746
|
+
if not tables:
|
|
747
|
+
raise FileNotFoundError(
|
|
748
|
+
f"detect_platform: {data_path} has no spatial/ directory and no "
|
|
749
|
+
f"counts table (.tsv/.csv/.txt, optionally gzipped)."
|
|
750
|
+
)
|
|
751
|
+
table = tables[0]
|
|
752
|
+
else:
|
|
753
|
+
table = data_path
|
|
754
|
+
|
|
755
|
+
labels = _peek_table_labels(table)
|
|
756
|
+
if labels is None:
|
|
757
|
+
raise ValueError(
|
|
758
|
+
f"detect_platform: could not read coordinate labels from {table}. "
|
|
759
|
+
f"Pass platform= explicitly."
|
|
760
|
+
)
|
|
761
|
+
parsed = _parse_coord_labels(labels)
|
|
762
|
+
if parsed is None:
|
|
763
|
+
raise ValueError(
|
|
764
|
+
f"detect_platform: labels in {table} are not coordinate-like "
|
|
765
|
+
f"(expected 'AxB' or 'A_B'). Pass platform= explicitly."
|
|
766
|
+
)
|
|
767
|
+
arr = np.asarray(parsed, dtype=float)
|
|
768
|
+
integral = np.allclose(arr, np.round(arr))
|
|
769
|
+
if integral and arr.min() >= 0 and arr.max() <= 50:
|
|
770
|
+
return 'dbit_seq'
|
|
771
|
+
raise ValueError(
|
|
772
|
+
f"detect_platform: coordinate labels in {table} do not match a known "
|
|
773
|
+
f"platform (integer indices within a 50 x 50 grid are expected for "
|
|
774
|
+
f"DBiT-seq). Pass platform= explicitly."
|
|
775
|
+
)
|
|
776
|
+
|
|
777
|
+
|
|
778
|
+
def _find_count_tables(directory: str) -> List[str]:
|
|
779
|
+
"""Delimited counts tables in a directory.
|
|
780
|
+
|
|
781
|
+
When several candidates exist the largest is taken, because the counts
|
|
782
|
+
matrix is invariably the biggest table in a sample directory.
|
|
783
|
+
"""
|
|
784
|
+
pats = ('*.tsv', '*.tsv.gz', '*.csv', '*.csv.gz', '*.txt', '*.txt.gz')
|
|
785
|
+
hits: List[str] = []
|
|
786
|
+
for p in pats:
|
|
787
|
+
hits.extend(glob.glob(os.path.join(directory, p)))
|
|
788
|
+
return sorted(hits, key=lambda h: -os.path.getsize(h))
|
|
789
|
+
|
|
790
|
+
|
|
791
|
+
def _read_counts_table(path: str, sep: str, chunk_rows: int = 2000):
|
|
792
|
+
"""Read a delimited dense counts table into a sparse matrix, in chunks.
|
|
793
|
+
|
|
794
|
+
Returns ``(csr, row_labels, col_labels)`` with rows and columns in file
|
|
795
|
+
order; the caller decides which axis is which.
|
|
796
|
+
|
|
797
|
+
Dense text matrices from these platforms can be large: a whole-transcriptome
|
|
798
|
+
grid is easily a gigabyte of text, which as a dense float64 frame would not
|
|
799
|
+
open on the 16 GB workstation this protocol targets. Parsing line by line
|
|
800
|
+
and keeping only the non-zeros holds peak memory at roughly one row, and the
|
|
801
|
+
result is a few hundred megabytes at most because these matrices are 95-99%
|
|
802
|
+
zeros. Files under 100 MB skip the machinery and are read in one go.
|
|
803
|
+
"""
|
|
804
|
+
size_mb = os.path.getsize(path) / 1e6
|
|
805
|
+
|
|
806
|
+
# Large text matrices are parsed once and cached next to the source, because
|
|
807
|
+
# streaming 1.8 GB of text takes minutes while reloading the compressed
|
|
808
|
+
# sparse form takes under a second. The cache is invalidated whenever the
|
|
809
|
+
# source file is newer.
|
|
810
|
+
cache = path + '.codeconv_cache.npz'
|
|
811
|
+
if size_mb >= 100 and os.path.exists(cache) and \
|
|
812
|
+
os.path.getmtime(cache) >= os.path.getmtime(path):
|
|
813
|
+
z = np.load(cache, allow_pickle=False)
|
|
814
|
+
mat = sp.csr_matrix((z['data'], z['indices'], z['indptr']), shape=tuple(z['shape']))
|
|
815
|
+
print(f" cached parse: {cache}")
|
|
816
|
+
return mat, [str(x) for x in z['rows']], [str(x) for x in z['cols']]
|
|
817
|
+
|
|
818
|
+
if size_mb < 100:
|
|
819
|
+
df = pd.read_csv(path, sep=sep, index_col=0)
|
|
820
|
+
if df.shape[1] == 0: # wrong delimiter guess
|
|
821
|
+
df = pd.read_csv(path, sep=None, engine='python', index_col=0)
|
|
822
|
+
# A trailing delimiter on the header line yields a blank all-NaN column.
|
|
823
|
+
df = df.loc[:, [c for c in df.columns if str(c).strip() != '' and not df[c].isna().all()]]
|
|
824
|
+
return sp.csr_matrix(df.to_numpy(dtype=np.float64)), \
|
|
825
|
+
[str(i) for i in df.index], [str(c) for c in df.columns]
|
|
826
|
+
|
|
827
|
+
# Streaming line-by-line parse. pandas carries a large per-column overhead,
|
|
828
|
+
# which dominates on matrices tens of thousands of columns wide, so each
|
|
829
|
+
# line is parsed straight into a numeric array with
|
|
830
|
+
# np.fromstring and reduced to its non-zeros immediately. Peak memory is one
|
|
831
|
+
# row, and only the non-zeros are retained.
|
|
832
|
+
opener = gzip.open if path.endswith('.gz') else open
|
|
833
|
+
with opener(path, 'rt') as f:
|
|
834
|
+
header = f.readline().rstrip('\r\n').split(sep)
|
|
835
|
+
col_labels = [str(c) for c in header[1:] if str(c).strip() != '']
|
|
836
|
+
n_cols = len(col_labels)
|
|
837
|
+
print(f" large table ({size_mb:.0f} MB, {n_cols} columns); streaming")
|
|
838
|
+
|
|
839
|
+
indptr = [0]
|
|
840
|
+
indices: List[np.ndarray] = []
|
|
841
|
+
values: List[np.ndarray] = []
|
|
842
|
+
row_labels: List[str] = []
|
|
843
|
+
nnz = 0
|
|
844
|
+
for line in f:
|
|
845
|
+
label, _, rest = line.rstrip('\r\n').partition(sep)
|
|
846
|
+
if not label:
|
|
847
|
+
continue
|
|
848
|
+
row = np.fromstring(rest, dtype=np.float32, sep=sep)
|
|
849
|
+
if row.size > n_cols:
|
|
850
|
+
row = row[:n_cols]
|
|
851
|
+
elif row.size < n_cols:
|
|
852
|
+
row = np.pad(row, (0, n_cols - row.size))
|
|
853
|
+
nz = np.flatnonzero(row)
|
|
854
|
+
indices.append(nz.astype(np.int32))
|
|
855
|
+
values.append(row[nz])
|
|
856
|
+
nnz += nz.size
|
|
857
|
+
indptr.append(nnz)
|
|
858
|
+
row_labels.append(str(label))
|
|
859
|
+
if len(row_labels) % 2000 == 0:
|
|
860
|
+
print(f" {len(row_labels)} rows, {nnz:,} non-zeros", end='\r')
|
|
861
|
+
|
|
862
|
+
print(f" {len(row_labels)} rows, {nnz:,} non-zeros ")
|
|
863
|
+
mat = sp.csr_matrix(
|
|
864
|
+
(np.concatenate(values) if values else np.zeros(0, np.float32),
|
|
865
|
+
np.concatenate(indices) if indices else np.zeros(0, np.int32),
|
|
866
|
+
np.asarray(indptr, dtype=np.int64)),
|
|
867
|
+
shape=(len(row_labels), n_cols),
|
|
868
|
+
)
|
|
869
|
+
try:
|
|
870
|
+
np.savez_compressed(
|
|
871
|
+
cache, data=mat.data, indices=mat.indices, indptr=mat.indptr,
|
|
872
|
+
shape=np.asarray(mat.shape), rows=np.asarray(row_labels),
|
|
873
|
+
cols=np.asarray(col_labels),
|
|
874
|
+
)
|
|
875
|
+
print(f" cached the parse to {os.path.basename(cache)} for future runs")
|
|
876
|
+
except Exception as exc: # a read-only data directory must not be fatal
|
|
877
|
+
print(f" (could not write parse cache: {exc})")
|
|
878
|
+
return mat, row_labels, col_labels
|
|
879
|
+
|
|
880
|
+
|
|
881
|
+
def _dbit_pixel_um(name: str, default: float = 50.0) -> float:
|
|
882
|
+
"""Infer the DBiT-seq microchannel width from a sample name.
|
|
883
|
+
|
|
884
|
+
GEO sample names in the DBiT-seq series encode the resolution, e.g.
|
|
885
|
+
``GSM4189611_50t`` is the 50 um tail section and ``GSM4189615`` variants
|
|
886
|
+
carry ``10``. Returns the default when the name says nothing.
|
|
887
|
+
"""
|
|
888
|
+
m = re.search(r'(?<!\d)(10|25|50)(?=[^\d]|$)', str(name))
|
|
889
|
+
return float(m.group(1)) if m else float(default)
|
|
890
|
+
|
|
891
|
+
|
|
892
|
+
def _peek_table_labels(path: str, n: int = 200) -> Optional[List[str]]:
|
|
893
|
+
"""First-column labels of a delimited table, without reading the whole file."""
|
|
894
|
+
base = os.path.basename(path).lower()
|
|
895
|
+
try:
|
|
896
|
+
sep = ',' if '.csv' in base else '\t'
|
|
897
|
+
head = pd.read_csv(path, sep=sep, index_col=0, nrows=n)
|
|
898
|
+
if head.shape[1] == 0: # wrong delimiter guess, e.g. a comma-separated .txt
|
|
899
|
+
head = pd.read_csv(path, sep=None, engine='python', index_col=0, nrows=n)
|
|
900
|
+
except Exception:
|
|
901
|
+
return None
|
|
902
|
+
labels = [str(x) for x in head.index]
|
|
903
|
+
# Coordinate labels may live in the header instead (genes x spots orientation).
|
|
904
|
+
if _parse_coord_labels(labels) is None:
|
|
905
|
+
cols = [str(c) for c in head.columns]
|
|
906
|
+
if _parse_coord_labels(cols) is not None:
|
|
907
|
+
return cols
|
|
908
|
+
return labels
|
|
909
|
+
|
|
910
|
+
|
|
911
|
+
def _parse_coord_labels(labels) -> Optional[np.ndarray]:
|
|
912
|
+
"""Parse ``AxB`` / ``A_B`` / ``A,B`` labels into an (n, 2) float array, or None."""
|
|
913
|
+
pat = re.compile(r'^\s*(-?\d+(?:\.\d+)?)\s*[x_,\-]\s*(-?\d+(?:\.\d+)?)\s*$', re.I)
|
|
914
|
+
out = []
|
|
915
|
+
for lab in labels:
|
|
916
|
+
m = pat.match(str(lab))
|
|
917
|
+
if m is None:
|
|
918
|
+
return None
|
|
919
|
+
out.append((float(m.group(1)), float(m.group(2))))
|
|
920
|
+
return np.asarray(out, dtype=float) if out else None
|
|
921
|
+
|
|
922
|
+
|
|
923
|
+
def make_dummy_tissue_image(
|
|
924
|
+
coords_fullres: np.ndarray,
|
|
925
|
+
spatial_dir: str,
|
|
926
|
+
spot_diameter_fullres: Optional[float] = None,
|
|
927
|
+
margin_frac: float = 0.03,
|
|
928
|
+
hires_px: int = 2000,
|
|
929
|
+
lowres_px: int = 600,
|
|
930
|
+
quiet: bool = False,
|
|
931
|
+
) -> dict:
|
|
932
|
+
"""Write a black placeholder tissue image and matching scale factors.
|
|
933
|
+
|
|
934
|
+
Several spot-based platforms ship no registered histology at all (DBiT-seq
|
|
935
|
+
samples vary), but ``Seurat::Load10X_Spatial`` and the rest of the 10x ecosystem
|
|
936
|
+
expect an image to exist. Rather than fail, write a black canvas sized to the
|
|
937
|
+
spot grid so the object loads and every spatial plot still works; the tissue
|
|
938
|
+
background is simply blank.
|
|
939
|
+
|
|
940
|
+
The canvas is anchored at the coordinate origin and extended to the maximum
|
|
941
|
+
coordinate plus a small margin, so the existing pixel coordinates stay valid
|
|
942
|
+
and no barcode-to-position mapping has to be rewritten.
|
|
943
|
+
|
|
944
|
+
Parameters
|
|
945
|
+
----------
|
|
946
|
+
coords_fullres : np.ndarray
|
|
947
|
+
(n, >=2) array of full-resolution pixel coordinates as (row, col).
|
|
948
|
+
spatial_dir : str
|
|
949
|
+
Directory to write ``tissue_hires_image.png``, ``tissue_lowres_image.png``
|
|
950
|
+
into. Created if missing.
|
|
951
|
+
spot_diameter_fullres : Optional[float]
|
|
952
|
+
Spot footprint in full-resolution pixels, for the returned scale factors.
|
|
953
|
+
Defaults to 1/40th of the larger canvas dimension.
|
|
954
|
+
margin_frac : float
|
|
955
|
+
Blank margin beyond the outermost spot, as a fraction of the extent.
|
|
956
|
+
|
|
957
|
+
Returns
|
|
958
|
+
-------
|
|
959
|
+
dict
|
|
960
|
+
Scale-factor entries (``tissue_hires_scalef``, ``tissue_lowres_scalef``,
|
|
961
|
+
``spot_diameter_fullres``, ``fiducial_diameter_fullres``) to write into
|
|
962
|
+
``scalefactors_json.json``.
|
|
963
|
+
"""
|
|
964
|
+
import matplotlib.image as mpimg
|
|
965
|
+
|
|
966
|
+
xy = np.asarray(coords_fullres, dtype=float)[:, :2]
|
|
967
|
+
if xy.size == 0:
|
|
968
|
+
raise ValueError("make_dummy_tissue_image: no coordinates given.")
|
|
969
|
+
|
|
970
|
+
extent = np.nanmax(xy, axis=0)
|
|
971
|
+
pad = margin_frac * float(np.nanmax(extent)) if np.nanmax(extent) > 0 else 1.0
|
|
972
|
+
h_full = float(extent[0] + pad)
|
|
973
|
+
w_full = float(extent[1] + pad)
|
|
974
|
+
if not (np.isfinite(h_full) and np.isfinite(w_full)) or min(h_full, w_full) <= 0:
|
|
975
|
+
raise ValueError("make_dummy_tissue_image: degenerate coordinate extent.")
|
|
976
|
+
|
|
977
|
+
if spot_diameter_fullres is None:
|
|
978
|
+
spot_diameter_fullres = max(h_full, w_full) / 40.0
|
|
979
|
+
|
|
980
|
+
os.makedirs(spatial_dir, exist_ok=True)
|
|
981
|
+
scalefs = {}
|
|
982
|
+
for fname, target in (("tissue_hires_image.png", hires_px),
|
|
983
|
+
("tissue_lowres_image.png", lowres_px)):
|
|
984
|
+
scalef = float(target) / max(h_full, w_full)
|
|
985
|
+
h_px = max(1, int(round(h_full * scalef)))
|
|
986
|
+
w_px = max(1, int(round(w_full * scalef)))
|
|
987
|
+
mpimg.imsave(os.path.join(spatial_dir, fname),
|
|
988
|
+
np.zeros((h_px, w_px, 3), dtype=np.uint8))
|
|
989
|
+
key = 'tissue_hires_scalef' if 'hires' in fname else 'tissue_lowres_scalef'
|
|
990
|
+
scalefs[key] = scalef
|
|
991
|
+
|
|
992
|
+
scalefs['spot_diameter_fullres'] = float(spot_diameter_fullres)
|
|
993
|
+
scalefs['fiducial_diameter_fullres'] = float(spot_diameter_fullres) * 1.75
|
|
994
|
+
|
|
995
|
+
if not quiet:
|
|
996
|
+
print(
|
|
997
|
+
f" WARNING: no tissue image detected, providing a dummy one. "
|
|
998
|
+
f"Wrote a black {int(w_full)} x {int(h_full)} full-resolution canvas "
|
|
999
|
+
f"({margin_frac:.0%} margin) covering {xy.shape[0]} spots. Spatial plots will "
|
|
1000
|
+
f"render on a blank background; H&E overlays are not available for this sample."
|
|
1001
|
+
)
|
|
1002
|
+
return scalefs
|
|
1003
|
+
|
|
1004
|
+
|
|
1005
|
+
# --- Built-in species configuration -------------------------------------------
|
|
1006
|
+
#
|
|
1007
|
+
# The human and mouse profiles ship inside the module so a fresh install runs
|
|
1008
|
+
# without any extra download: pass config_path=None (or point it at a file that
|
|
1009
|
+
# does not exist) and these are used. An external codeconv_config.json always
|
|
1010
|
+
# wins when it is present, so existing runs and hand-tuned profiles are
|
|
1011
|
+
# unaffected. Regenerate a species block with the CoexpressDeconvolve Config
|
|
1012
|
+
# Maker and write it to a JSON file to override.
|
|
1013
|
+
#
|
|
1014
|
+
# NOTE: the mouse hk_profiles and engine_parameters below mirror the human
|
|
1015
|
+
# reference and are placeholders. Recalibrate them from a mouse scRNA-seq
|
|
1016
|
+
# reference before publication-grade mouse runs. The 'other' block is
|
|
1017
|
+
# intentionally empty: non-hs/mm organisms must supply their own profile.
|
|
1018
|
+
DEFAULT_CONFIG: dict = {
|
|
1019
|
+
"min_topic_percentage": 0.05,
|
|
1020
|
+
"species_profiles": {
|
|
1021
|
+
"hs": {
|
|
1022
|
+
"noise_regex": "^MT-|^RP[SL][0-9]+|^LINC|^MIR|^AC[0-9]+",
|
|
1023
|
+
"hk_profiles": {
|
|
1024
|
+
"RPL19": 2.529822467672188, "TUBB": 0.6653868555153559,
|
|
1025
|
+
"EEF1G": 1.409845725236883, "PPIA": 1.4689432838260164,
|
|
1026
|
+
"GAPDH": 1.9870339974714675, "ABCF1": 0.2683943094020617,
|
|
1027
|
+
"SDHA": 0.22436799098809676, "OAZ1": 1.5212192863178522,
|
|
1028
|
+
"G6PD": 0.10922542882001043, "ALAS1": 0.10135028595398016,
|
|
1029
|
+
"GUSB": 0.1690932556829679, "HPRT1": 0.18642984163014426,
|
|
1030
|
+
"POLR2A": 0.36647764940764155, "POLR1B": 0.040193901737717995,
|
|
1031
|
+
"TBP": 0.060567363961918315,
|
|
1032
|
+
},
|
|
1033
|
+
"qc_markers": [
|
|
1034
|
+
"CD3E", "CD4", "CD8A", "MS4A1", "CD19", "PTPRC",
|
|
1035
|
+
"HBB", "CD14", "FCGR3A", "CD34", "NCAM1", "JCHAIN",
|
|
1036
|
+
"EPCAM", "KRT18", "COL1A1", "DCN", "PECAM1", "ERBB2",
|
|
1037
|
+
],
|
|
1038
|
+
"engine_parameters": {"mu": 5483.714684902277, "phi": 0.7002926708915271},
|
|
1039
|
+
},
|
|
1040
|
+
"mm": {
|
|
1041
|
+
"noise_regex": "^mt-|^Rp[sl][0-9]+|^Gm[0-9]+|^Mir|^Rik$",
|
|
1042
|
+
"hk_profiles": {
|
|
1043
|
+
"Rpl19": 2.529822467672188, "Tubb5": 0.6653868555153559,
|
|
1044
|
+
"Eef1g": 1.409845725236883, "Ppia": 1.4689432838260164,
|
|
1045
|
+
"Gapdh": 1.9870339974714675, "Abcf1": 0.2683943094020617,
|
|
1046
|
+
"Sdha": 0.22436799098809676, "Oaz1": 1.5212192863178522,
|
|
1047
|
+
"G6pdx": 0.10922542882001043, "Alas1": 0.10135028595398016,
|
|
1048
|
+
"Gusb": 0.1690932556829679, "Hprt": 0.18642984163014426,
|
|
1049
|
+
"Polr2a": 0.36647764940764155, "Polr1b": 0.040193901737717995,
|
|
1050
|
+
"Tbp": 0.060567363961918315,
|
|
1051
|
+
},
|
|
1052
|
+
"qc_markers": [
|
|
1053
|
+
"Cd3e", "Cd4", "Cd8a", "Ms4a1", "Cd19", "Ptprc",
|
|
1054
|
+
"Hbb-bs", "Cd14", "Fcgr3", "Cd34", "Ncam1", "Jchain",
|
|
1055
|
+
"Epcam", "Krt18", "Col1a1", "Dcn", "Pecam1", "Erbb2",
|
|
1056
|
+
],
|
|
1057
|
+
"engine_parameters": {"mu": 5483.714684902277, "phi": 0.7002926708915271},
|
|
1058
|
+
},
|
|
1059
|
+
"other": {
|
|
1060
|
+
"noise_regex": "",
|
|
1061
|
+
"hk_profiles": {},
|
|
1062
|
+
"qc_markers": [],
|
|
1063
|
+
"engine_parameters": {"mu": 5000.0, "phi": 0.7},
|
|
1064
|
+
},
|
|
1065
|
+
},
|
|
1066
|
+
}
|
|
1067
|
+
|
|
1068
|
+
|
|
1069
|
+
def write_default_config(path: str = "codeconv_config.json", overwrite: bool = False) -> str:
|
|
1070
|
+
"""Write the built-in species configuration to a JSON file, for editing.
|
|
1071
|
+
|
|
1072
|
+
Use this when you want to hand-tune a profile or add an organism: dump the
|
|
1073
|
+
defaults, edit the file, then pass its path as ``config_path``.
|
|
1074
|
+
"""
|
|
1075
|
+
if os.path.exists(path) and not overwrite:
|
|
1076
|
+
raise FileExistsError(f"{path} already exists; pass overwrite=True to replace it.")
|
|
1077
|
+
with open(path, 'w') as f:
|
|
1078
|
+
json.dump(DEFAULT_CONFIG, f, indent=4)
|
|
1079
|
+
print(f"Wrote built-in species configuration to {path}")
|
|
1080
|
+
return path
|
|
1081
|
+
|
|
1082
|
+
|
|
1083
|
+
# Helpers
|
|
1084
|
+
|
|
1085
|
+
def _load_config(config_path: Optional[str], species: str) -> dict:
|
|
1086
|
+
"""Load a species profile, falling back to the built-in configuration.
|
|
1087
|
+
|
|
1088
|
+
An external JSON at ``config_path`` wins when it exists. Passing None, or a
|
|
1089
|
+
path that is not there, uses ``DEFAULT_CONFIG``, so a fresh install runs
|
|
1090
|
+
without downloading anything.
|
|
1091
|
+
"""
|
|
1092
|
+
import copy
|
|
1093
|
+
|
|
1094
|
+
if config_path and os.path.exists(config_path):
|
|
1095
|
+
with open(config_path, 'r') as f:
|
|
1096
|
+
cfg = json.load(f)
|
|
1097
|
+
else:
|
|
1098
|
+
if config_path:
|
|
1099
|
+
print(f" config '{config_path}' not found; using the built-in species profiles.")
|
|
1100
|
+
cfg = copy.deepcopy(DEFAULT_CONFIG)
|
|
1101
|
+
|
|
1102
|
+
if species not in cfg["species_profiles"]:
|
|
1103
|
+
raise ValueError(
|
|
1104
|
+
f"Species '{species}' not in config. Available: {list(cfg['species_profiles'].keys())}"
|
|
1105
|
+
)
|
|
1106
|
+
profile = copy.deepcopy(cfg["species_profiles"][species])
|
|
1107
|
+
if not profile.get("hk_profiles"):
|
|
1108
|
+
raise ValueError(
|
|
1109
|
+
f"Species '{species}' has an empty hk_profiles block, so cell-density calibration "
|
|
1110
|
+
f"has nothing to calibrate against. Built-in profiles exist for 'hs' and 'mm'; for "
|
|
1111
|
+
f"any other organism, generate a species block with the CoexpressDeconvolve Config "
|
|
1112
|
+
f"Maker ('Estimate single cell parameters.ipynb'), save it with write_default_config() "
|
|
1113
|
+
f"as a starting point, and pass that file as config_path."
|
|
1114
|
+
)
|
|
1115
|
+
profile["min_topic_percentage"] = cfg.get("min_topic_percentage", 0.05)
|
|
1116
|
+
profile["species"] = species
|
|
1117
|
+
return profile
|
|
1118
|
+
|
|
1119
|
+
|
|
1120
|
+
def _normalize_paths(spatial_path) -> Dict[str, str]:
|
|
1121
|
+
"""Coerce a string / list / dict input into a {name: path} dict.
|
|
1122
|
+
|
|
1123
|
+
Single string -> {basename: path}.
|
|
1124
|
+
List of strings -> {basename(p): p for p in list}; duplicate basenames are an error.
|
|
1125
|
+
Dict -> passthrough.
|
|
1126
|
+
"""
|
|
1127
|
+
if isinstance(spatial_path, str):
|
|
1128
|
+
name = os.path.basename(spatial_path.rstrip('/')) or spatial_path
|
|
1129
|
+
return {name: spatial_path}
|
|
1130
|
+
if isinstance(spatial_path, dict):
|
|
1131
|
+
return dict(spatial_path)
|
|
1132
|
+
if isinstance(spatial_path, (list, tuple)):
|
|
1133
|
+
names = [os.path.basename(p.rstrip('/')) or p for p in spatial_path]
|
|
1134
|
+
if len(set(names)) != len(names):
|
|
1135
|
+
raise ValueError(
|
|
1136
|
+
"Duplicate slice basenames detected in spatial_path list; "
|
|
1137
|
+
"use the dict form to disambiguate: {name: path, ...}."
|
|
1138
|
+
)
|
|
1139
|
+
return dict(zip(names, spatial_path))
|
|
1140
|
+
raise TypeError(f"spatial_path must be str, list, or dict, got {type(spatial_path)}")
|
|
1141
|
+
|
|
1142
|
+
|
|
1143
|
+
def _broadcast(param, slice_names: List[str], default=None):
|
|
1144
|
+
"""Expand a scalar/dict param into a {name: value} dict aligned with slice_names."""
|
|
1145
|
+
if isinstance(param, dict):
|
|
1146
|
+
for k in slice_names:
|
|
1147
|
+
if k not in param:
|
|
1148
|
+
if default is not None:
|
|
1149
|
+
param[k] = default
|
|
1150
|
+
else:
|
|
1151
|
+
raise KeyError(f"Per-slice param missing entry for slice '{k}'.")
|
|
1152
|
+
return param
|
|
1153
|
+
if param is None:
|
|
1154
|
+
return {n: default for n in slice_names}
|
|
1155
|
+
return {n: param for n in slice_names}
|
|
1156
|
+
|
|
1157
|
+
|
|
1158
|
+
def _variational_e_step(
|
|
1159
|
+
X: sp.csr_matrix,
|
|
1160
|
+
beta: np.ndarray,
|
|
1161
|
+
alpha: float,
|
|
1162
|
+
max_iter: int = 50,
|
|
1163
|
+
tol: float = 1e-3,
|
|
1164
|
+
) -> np.ndarray:
|
|
1165
|
+
"""Estimate theta given fixed beta via mean-field variational LDA inference.
|
|
1166
|
+
|
|
1167
|
+
Standard textbook update: gamma_d = alpha + sum_n count_n * phi_{n,k}
|
|
1168
|
+
with phi_{n,k} proportional to beta[k, w_n] * exp(digamma(gamma_d[k])).
|
|
1169
|
+
|
|
1170
|
+
X: (n_docs, n_words) sparse counts.
|
|
1171
|
+
beta: (K, n_words), rows sum to 1.
|
|
1172
|
+
alpha: scalar Dirichlet prior on theta.
|
|
1173
|
+
"""
|
|
1174
|
+
X_csr = X.tocsr()
|
|
1175
|
+
n_docs, n_words = X_csr.shape
|
|
1176
|
+
K = beta.shape[0]
|
|
1177
|
+
|
|
1178
|
+
doc_lens = np.array(X_csr.sum(axis=1)).flatten()
|
|
1179
|
+
gamma = np.full((n_docs, K), float(alpha)) + (doc_lens[:, None] / K)
|
|
1180
|
+
|
|
1181
|
+
for it in range(max_iter):
|
|
1182
|
+
gamma_old = gamma.copy()
|
|
1183
|
+
Elogtheta = digamma(gamma) - digamma(gamma.sum(axis=1, keepdims=True))
|
|
1184
|
+
expElogtheta = np.exp(Elogtheta)
|
|
1185
|
+
|
|
1186
|
+
new_gamma = np.full((n_docs, K), float(alpha))
|
|
1187
|
+
for d in range(n_docs):
|
|
1188
|
+
start, end = X_csr.indptr[d], X_csr.indptr[d + 1]
|
|
1189
|
+
if start == end:
|
|
1190
|
+
continue
|
|
1191
|
+
cols = X_csr.indices[start:end]
|
|
1192
|
+
vals = X_csr.data[start:end].astype(float)
|
|
1193
|
+
|
|
1194
|
+
phinorm = expElogtheta[d] @ beta[:, cols]
|
|
1195
|
+
phinorm = np.maximum(phinorm, 1e-100)
|
|
1196
|
+
ratio = vals / phinorm
|
|
1197
|
+
contrib = expElogtheta[d] * (beta[:, cols] @ ratio)
|
|
1198
|
+
new_gamma[d] = alpha + contrib
|
|
1199
|
+
|
|
1200
|
+
gamma = new_gamma
|
|
1201
|
+
if np.mean(np.abs(gamma - gamma_old)) < tol:
|
|
1202
|
+
break
|
|
1203
|
+
|
|
1204
|
+
theta = gamma / gamma.sum(axis=1, keepdims=True)
|
|
1205
|
+
return theta
|
|
1206
|
+
|
|
1207
|
+
|
|
1208
|
+
def _align_topics(betas: Dict[str, np.ndarray], anchor: Optional[str] = None) -> Dict[str, np.ndarray]:
|
|
1209
|
+
"""Reorder per-slice topics so that index k means the same biological program in every slice.
|
|
1210
|
+
|
|
1211
|
+
Hungarian matching on cosine similarity of beta rows against an anchor slice.
|
|
1212
|
+
Returns a dict of betas with rows reordered to anchor's topic order.
|
|
1213
|
+
"""
|
|
1214
|
+
names = list(betas.keys())
|
|
1215
|
+
if anchor is None:
|
|
1216
|
+
anchor = names[0]
|
|
1217
|
+
anchor_beta = betas[anchor]
|
|
1218
|
+
|
|
1219
|
+
def _row_normalize(M):
|
|
1220
|
+
norms = np.linalg.norm(M, axis=1, keepdims=True)
|
|
1221
|
+
norms[norms == 0] = 1.0
|
|
1222
|
+
return M / norms
|
|
1223
|
+
|
|
1224
|
+
anchor_norm = _row_normalize(anchor_beta)
|
|
1225
|
+
aligned = {anchor: anchor_beta}
|
|
1226
|
+
for s in names:
|
|
1227
|
+
if s == anchor:
|
|
1228
|
+
continue
|
|
1229
|
+
b = betas[s]
|
|
1230
|
+
b_norm = _row_normalize(b)
|
|
1231
|
+
sim = anchor_norm @ b_norm.T
|
|
1232
|
+
# Hungarian minimizes; we want to maximize similarity, so negate.
|
|
1233
|
+
row_ind, col_ind = linear_sum_assignment(-sim)
|
|
1234
|
+
new_b = np.zeros_like(b)
|
|
1235
|
+
new_b[row_ind] = b[col_ind]
|
|
1236
|
+
aligned[s] = new_b
|
|
1237
|
+
return aligned
|
|
1238
|
+
|
|
1239
|
+
|
|
1240
|
+
def _safe_multinomial(rng, n, p):
|
|
1241
|
+
"""multinomial sample that tolerates float roundoff.
|
|
1242
|
+
|
|
1243
|
+
Tries a direct rng.multinomial(n, p) call. Falls back to a clip + normalize + shave
|
|
1244
|
+
pass only when vanilla raises ValueError (numpy strictly checks
|
|
1245
|
+
pvals[:-1].sum() > 1.0 and rejects ULP-level overflow).
|
|
1246
|
+
"""
|
|
1247
|
+
p = np.asarray(p, dtype=np.float64)
|
|
1248
|
+
try:
|
|
1249
|
+
return rng.multinomial(int(n), p)
|
|
1250
|
+
except ValueError:
|
|
1251
|
+
pass
|
|
1252
|
+
p = np.clip(p, 0.0, None)
|
|
1253
|
+
s = p.sum()
|
|
1254
|
+
if s <= 0 or not np.isfinite(s):
|
|
1255
|
+
out = np.zeros(len(p), dtype=int)
|
|
1256
|
+
out[0] = int(n)
|
|
1257
|
+
return out
|
|
1258
|
+
p = p / s
|
|
1259
|
+
overflow = p[:-1].sum() - 1.0
|
|
1260
|
+
if overflow > 0:
|
|
1261
|
+
idx = int(np.argmax(p[:-1]))
|
|
1262
|
+
p[idx] = max(0.0, p[idx] - overflow)
|
|
1263
|
+
return rng.multinomial(int(n), p)
|
|
1264
|
+
|
|
1265
|
+
|
|
1266
|
+
def _inv_digamma(y, n_iter: int = 5):
|
|
1267
|
+
"""Newton's method for the inverse of the digamma function.
|
|
1268
|
+
|
|
1269
|
+
Initialization rule from Minka, "Estimating a Dirichlet distribution" (2000),
|
|
1270
|
+
Appendix C. Five Newton steps suffice for double precision.
|
|
1271
|
+
"""
|
|
1272
|
+
y = np.asarray(y, dtype=float)
|
|
1273
|
+
x = np.where(y >= -2.22, np.exp(y) + 0.5, -1.0 / (y - digamma(1.0)))
|
|
1274
|
+
for _ in range(n_iter):
|
|
1275
|
+
x = x - (digamma(x) - y) / polygamma(1, x)
|
|
1276
|
+
return x
|
|
1277
|
+
|
|
1278
|
+
|
|
1279
|
+
def _dirichlet_alpha_mle(
|
|
1280
|
+
theta: np.ndarray,
|
|
1281
|
+
max_iter: int = 200,
|
|
1282
|
+
tol: float = 1e-7,
|
|
1283
|
+
eps: float = 1e-12,
|
|
1284
|
+
) -> np.ndarray:
|
|
1285
|
+
"""Estimate an asymmetric Dirichlet alpha from observed proportions.
|
|
1286
|
+
|
|
1287
|
+
Implements Minka's fixed-point iteration (2000, eq. 9). Treats each row of
|
|
1288
|
+
theta as an observed sample from Dir(alpha). The iteration loops
|
|
1289
|
+
alpha_k <- digamma^{-1}( digamma(sum_k alpha_k) + E_d[log theta_{d,k}] )
|
|
1290
|
+
until convergence.
|
|
1291
|
+
|
|
1292
|
+
Used in Step 5 as a post-hoc proxy for the alpha parameter that
|
|
1293
|
+
STdeconvolve's R LDA estimates directly (sklearn fixes doc_topic_prior, so
|
|
1294
|
+
we recover an alpha estimate from the fitted topic distributions instead).
|
|
1295
|
+
A fitted mean alpha < 1 indicates the model retained a sparse Dirichlet
|
|
1296
|
+
prior (each spot dominated by a few topics); values >= 1 mean topics smear
|
|
1297
|
+
across spots, which is the classical "K is too large" signature.
|
|
1298
|
+
|
|
1299
|
+
Parameters
|
|
1300
|
+
----------
|
|
1301
|
+
theta : (N, K) array of proportions, each row should sum to ~1.
|
|
1302
|
+
|
|
1303
|
+
Returns
|
|
1304
|
+
-------
|
|
1305
|
+
alpha : (K,) array of per-topic concentration parameters.
|
|
1306
|
+
"""
|
|
1307
|
+
theta = np.asarray(theta, dtype=float)
|
|
1308
|
+
theta = np.clip(theta, eps, 1.0)
|
|
1309
|
+
log_p_mean = np.mean(np.log(theta), axis=0)
|
|
1310
|
+
|
|
1311
|
+
# Method-of-moments initialization (Minka eq. 23).
|
|
1312
|
+
p_mean = np.mean(theta, axis=0)
|
|
1313
|
+
p_var = np.maximum(np.mean(theta ** 2, axis=0) - p_mean ** 2, eps)
|
|
1314
|
+
s_per_k = (p_mean * (1.0 - p_mean) / p_var) - 1.0
|
|
1315
|
+
s = max(float(np.median(s_per_k)), 0.1)
|
|
1316
|
+
alpha = np.maximum(p_mean * s, 1e-3)
|
|
1317
|
+
|
|
1318
|
+
for _ in range(max_iter):
|
|
1319
|
+
alpha_old = alpha.copy()
|
|
1320
|
+
alpha = _inv_digamma(digamma(alpha.sum()) + log_p_mean)
|
|
1321
|
+
alpha = np.maximum(alpha, 1e-6)
|
|
1322
|
+
if np.max(np.abs(alpha - alpha_old)) < tol:
|
|
1323
|
+
break
|
|
1324
|
+
return alpha
|
|
1325
|
+
|
|
1326
|
+
|
|
1327
|
+
def _overdispersed_genes(
|
|
1328
|
+
counts: sp.csr_matrix,
|
|
1329
|
+
n_top: int,
|
|
1330
|
+
poly_deg: int = 3,
|
|
1331
|
+
fit_mask: Optional[np.ndarray] = None,
|
|
1332
|
+
):
|
|
1333
|
+
"""Select overdispersed genes via mean-variance trend residuals.
|
|
1334
|
+
|
|
1335
|
+
Library-size-normalize counts to 10k, log1p-transform, compute per-gene
|
|
1336
|
+
mean and variance, fit a smoothed polynomial trend log10(var) ~
|
|
1337
|
+
poly(log10(mean)) across genes, and rank genes by the residual above the
|
|
1338
|
+
trend. The top n_top residuals are the overdispersed gene set.
|
|
1339
|
+
|
|
1340
|
+
This mirrors STdeconvolve's restrictCorpus / getOverdispersedGenes: genes
|
|
1341
|
+
whose variance exceeds the global mean-variance relationship are the ones
|
|
1342
|
+
that carry biological signal beyond Poisson sampling noise. Replaces the
|
|
1343
|
+
earlier binned-dispersion-z-score selection inspired by Seurat: more robust
|
|
1344
|
+
to bin-edge effects and closer to the corpus-building procedure LDA-based
|
|
1345
|
+
spatial deconvolution expects.
|
|
1346
|
+
|
|
1347
|
+
Parameters
|
|
1348
|
+
----------
|
|
1349
|
+
fit_mask : Optional[np.ndarray]
|
|
1350
|
+
Boolean mask of length n_genes. If provided, the polynomial trend is
|
|
1351
|
+
fit using only the genes where fit_mask is True (e.g. to exclude
|
|
1352
|
+
pre-filtered rare or ubiquitous genes), AND ranking is restricted to
|
|
1353
|
+
those same genes. Residuals are still computed for all genes so the
|
|
1354
|
+
diagnostic plot can show filtered-out genes as context.
|
|
1355
|
+
|
|
1356
|
+
Returns
|
|
1357
|
+
-------
|
|
1358
|
+
top_indices : np.ndarray (n_top_eff,)
|
|
1359
|
+
Gene indices sorted by residual, highest first. When fit_mask is given,
|
|
1360
|
+
all returned indices come from the masked-in set.
|
|
1361
|
+
log_mean, log_var, fitted_log_var, residuals : np.ndarray (n_genes,)
|
|
1362
|
+
Diagnostic arrays for plotting.
|
|
1363
|
+
"""
|
|
1364
|
+
row_sums = np.array(counts.sum(axis=1)).flatten()
|
|
1365
|
+
row_sums[row_sums == 0] = 1
|
|
1366
|
+
norm = counts.copy().astype(float)
|
|
1367
|
+
norm.data /= np.repeat(row_sums, np.diff(norm.indptr))
|
|
1368
|
+
norm.data *= 10000.0
|
|
1369
|
+
norm.data = np.log1p(norm.data)
|
|
1370
|
+
|
|
1371
|
+
mean_expr = np.array(norm.mean(axis=0)).flatten()
|
|
1372
|
+
sq = norm.copy()
|
|
1373
|
+
sq.data **= 2
|
|
1374
|
+
mean_sq = np.array(sq.mean(axis=0)).flatten()
|
|
1375
|
+
var_expr = np.maximum(mean_sq - mean_expr ** 2, 0.0)
|
|
1376
|
+
|
|
1377
|
+
eps = 1e-10
|
|
1378
|
+
log_mean = np.log10(mean_expr + eps)
|
|
1379
|
+
log_var = np.log10(var_expr + eps)
|
|
1380
|
+
valid = (mean_expr > 0) & (var_expr > 0)
|
|
1381
|
+
|
|
1382
|
+
if fit_mask is not None:
|
|
1383
|
+
fit_mask = np.asarray(fit_mask, dtype=bool)
|
|
1384
|
+
fit_valid = valid & fit_mask
|
|
1385
|
+
rank_pool = fit_valid
|
|
1386
|
+
else:
|
|
1387
|
+
fit_valid = valid
|
|
1388
|
+
rank_pool = valid
|
|
1389
|
+
|
|
1390
|
+
if int(fit_valid.sum()) < poly_deg + 1:
|
|
1391
|
+
# Pathological corpus: fall back to top by raw variance within the pool.
|
|
1392
|
+
n_top_eff = min(n_top, int(rank_pool.sum()) if rank_pool.sum() > 0 else n_top)
|
|
1393
|
+
scores = np.where(rank_pool, var_expr, -np.inf)
|
|
1394
|
+
top = np.argsort(scores)[-n_top_eff:][::-1]
|
|
1395
|
+
fitted = np.full_like(log_var, np.nan)
|
|
1396
|
+
return top, log_mean, log_var, fitted, np.zeros_like(log_var)
|
|
1397
|
+
|
|
1398
|
+
coeffs = np.polyfit(log_mean[fit_valid], log_var[fit_valid], deg=poly_deg)
|
|
1399
|
+
fitted_log_var = np.polyval(coeffs, log_mean)
|
|
1400
|
+
residuals = log_var - fitted_log_var
|
|
1401
|
+
residuals_for_ranking = np.where(rank_pool, residuals, -np.inf)
|
|
1402
|
+
|
|
1403
|
+
n_top_eff = min(n_top, int(rank_pool.sum()))
|
|
1404
|
+
top = np.argsort(residuals_for_ranking)[-n_top_eff:][::-1]
|
|
1405
|
+
return top, log_mean, log_var, fitted_log_var, residuals
|
|
1406
|
+
|
|
1407
|
+
|
|
1408
|
+
def _batched_multinomial(rng, n: np.ndarray, p: np.ndarray) -> np.ndarray:
|
|
1409
|
+
"""Batched multinomial sampler via sequential binomial decomposition.
|
|
1410
|
+
|
|
1411
|
+
Each row is independently drawn from Multinomial(n[i], p[i, :]). Works on
|
|
1412
|
+
both ``np.random.RandomState`` and the newer Generator
|
|
1413
|
+
types because both expose a vectorized ``binomial`` with broadcasting.
|
|
1414
|
+
|
|
1415
|
+
The recurrence is the standard composition: condition on the events
|
|
1416
|
+
already assigned and the remaining probability mass, draw the next column
|
|
1417
|
+
as a binomial of the remaining trials, then subtract.
|
|
1418
|
+
|
|
1419
|
+
Parameters
|
|
1420
|
+
----------
|
|
1421
|
+
rng : np.random.RandomState or np.random.Generator
|
|
1422
|
+
Source of randomness. Must support ``binomial(n, p)`` with vector args.
|
|
1423
|
+
n : (M,) array of nonnegative integer trial counts.
|
|
1424
|
+
p : (M, K) array of probability rows (each row should sum to ~1).
|
|
1425
|
+
|
|
1426
|
+
Returns
|
|
1427
|
+
-------
|
|
1428
|
+
(M, K) integer array. Rows sum exactly to ``n`` and entries are ``>= 0``.
|
|
1429
|
+
"""
|
|
1430
|
+
n = np.asarray(n, dtype=np.int64)
|
|
1431
|
+
p = np.asarray(p, dtype=np.float64)
|
|
1432
|
+
M, K = p.shape
|
|
1433
|
+
if M == 0 or K == 0:
|
|
1434
|
+
return np.zeros((M, K), dtype=np.int64)
|
|
1435
|
+
|
|
1436
|
+
result = np.zeros((M, K), dtype=np.int64)
|
|
1437
|
+
remaining = n.copy()
|
|
1438
|
+
# cum_p[:, k] = sum_{j >= k} p[:, j], i.e. mass still available at column k.
|
|
1439
|
+
cum_p = np.cumsum(p[:, ::-1], axis=1)[:, ::-1]
|
|
1440
|
+
for k in range(K - 1):
|
|
1441
|
+
denom = cum_p[:, k]
|
|
1442
|
+
# When denom is zero (all remaining mass collapses to zero) skip the
|
|
1443
|
+
# binomial and keep result[:, k] = 0.
|
|
1444
|
+
with np.errstate(invalid='ignore', divide='ignore'):
|
|
1445
|
+
prob_k = np.where(denom > 0, p[:, k] / denom, 0.0)
|
|
1446
|
+
prob_k = np.clip(prob_k, 0.0, 1.0)
|
|
1447
|
+
draw = rng.binomial(remaining, prob_k)
|
|
1448
|
+
draw = np.minimum(draw, remaining)
|
|
1449
|
+
result[:, k] = draw
|
|
1450
|
+
remaining = remaining - draw
|
|
1451
|
+
result[:, K - 1] = remaining
|
|
1452
|
+
return result
|
|
1453
|
+
|
|
1454
|
+
|
|
1455
|
+
def _topic_stability(aligned_betas: List[np.ndarray]) -> float:
|
|
1456
|
+
"""Mean pairwise cosine similarity of aligned topic rows across replicates.
|
|
1457
|
+
|
|
1458
|
+
Inputs are a list of (K, G) topic-gene matrices that have already been row-
|
|
1459
|
+
permuted into a common topic ordering (e.g. by Hungarian matching). For
|
|
1460
|
+
every pair of replicates, compute the per-topic cosine similarity and
|
|
1461
|
+
average across topics; then average across pairs.
|
|
1462
|
+
|
|
1463
|
+
Returns 1.0 for a single replicate (trivially stable) or perfectly
|
|
1464
|
+
identical replicates; values approach 0 as replicates diverge.
|
|
1465
|
+
"""
|
|
1466
|
+
n = len(aligned_betas)
|
|
1467
|
+
if n < 2:
|
|
1468
|
+
return 1.0
|
|
1469
|
+
sims = []
|
|
1470
|
+
for i in range(n):
|
|
1471
|
+
norm_i = np.linalg.norm(aligned_betas[i], axis=1)
|
|
1472
|
+
for j in range(i + 1, n):
|
|
1473
|
+
norm_j = np.linalg.norm(aligned_betas[j], axis=1)
|
|
1474
|
+
denom = np.maximum(norm_i * norm_j, 1e-12)
|
|
1475
|
+
cos = np.sum(aligned_betas[i] * aligned_betas[j], axis=1) / denom
|
|
1476
|
+
sims.append(float(np.mean(cos)))
|
|
1477
|
+
return float(np.mean(sims))
|
|
1478
|
+
|
|
1479
|
+
|
|
1480
|
+
def _topic_log2fc(beta: np.ndarray, eps: float = 1e-12) -> np.ndarray:
|
|
1481
|
+
"""Per-topic log2 fold change of beta vs the mean beta of all other topics.
|
|
1482
|
+
|
|
1483
|
+
log2fc[k, g] = log2( beta[k, g] / mean_{k' != k}( beta[k', g] ) )
|
|
1484
|
+
|
|
1485
|
+
Ranking genes by this quantity highlights what is *specific* to a topic
|
|
1486
|
+
rather than what is merely highly expressed everywhere, which is the
|
|
1487
|
+
interpretive lens STdeconvolve uses in its topGenes / getBetaTheta output.
|
|
1488
|
+
Falls through to zeros when there is only a single topic.
|
|
1489
|
+
|
|
1490
|
+
Parameters
|
|
1491
|
+
----------
|
|
1492
|
+
beta : (K, G) topic-gene distribution. Rows do not need to sum to 1 (we
|
|
1493
|
+
compare ratios, so any per-row scaling cancels).
|
|
1494
|
+
|
|
1495
|
+
Returns
|
|
1496
|
+
-------
|
|
1497
|
+
log2fc : (K, G) array.
|
|
1498
|
+
"""
|
|
1499
|
+
K = beta.shape[0]
|
|
1500
|
+
if K <= 1:
|
|
1501
|
+
return np.zeros_like(beta)
|
|
1502
|
+
total = beta.sum(axis=0, keepdims=True)
|
|
1503
|
+
mean_other = (total - beta) / (K - 1)
|
|
1504
|
+
return np.log2((beta + eps) / (mean_other + eps))
|
|
1505
|
+
|
|
1506
|
+
|
|
1507
|
+
# STEP 1: load
|
|
1508
|
+
|
|
1509
|
+
@_profile_step
|
|
1510
|
+
def step1_acquisition_and_anchoring(
|
|
1511
|
+
spatial_path=None,
|
|
1512
|
+
platform: str = "auto",
|
|
1513
|
+
visium_path=None,
|
|
1514
|
+
) -> Dict[str, SliceData]:
|
|
1515
|
+
"""Load spot-based spatial transcriptomics data per slice.
|
|
1516
|
+
|
|
1517
|
+
Accepts a str, list, or dict of paths, exactly as before. The platform is
|
|
1518
|
+
detected from the directory layout and announced; pass ``platform=`` to force
|
|
1519
|
+
it. Everything downstream is platform-agnostic, because Step 1 normalizes all
|
|
1520
|
+
inputs to the same four things: counts, coordinates, spot footprint, pitch.
|
|
1521
|
+
|
|
1522
|
+
Supported platforms (see :data:`PLATFORM_PROFILES` and :func:`detect_platform`):
|
|
1523
|
+
|
|
1524
|
+
============== ========================== ========================================
|
|
1525
|
+
``visium`` 55 um spots, 100 um pitch SpaceRanger layout, native, unchanged
|
|
1526
|
+
``dbit_seq`` 10/25/50 um, 2x the width grid-indexed table on a 50 x 50 array
|
|
1527
|
+
============== ========================== ========================================
|
|
1528
|
+
|
|
1529
|
+
Platforms whose capture units are already at or below the size of one cell -
|
|
1530
|
+
Slide-seq beads at 10 um, Visium HD 8 and 16 um bins, and the subcellular
|
|
1531
|
+
Stereo-seq bin1, Xenium and CosMx, are deliberately out of scope: a capture
|
|
1532
|
+
unit that holds at most one cell has nothing to deconvolve, and segmentation
|
|
1533
|
+
rather than a mixture model is the appropriate step there.
|
|
1534
|
+
|
|
1535
|
+
Platforms without histology get a black placeholder image and synthesized
|
|
1536
|
+
scale factors, so the exported object still loads in Seurat
|
|
1537
|
+
(see :func:`make_dummy_tissue_image`).
|
|
1538
|
+
|
|
1539
|
+
Parameters
|
|
1540
|
+
----------
|
|
1541
|
+
spatial_path : str | list | dict
|
|
1542
|
+
Dataset directory (or a {name: path} dict for multi-slice runs).
|
|
1543
|
+
platform : str
|
|
1544
|
+
``"auto"`` (default) or one of the keys of ``PLATFORM_PROFILES``. May also
|
|
1545
|
+
be a dict keyed by slice name for mixed-platform runs.
|
|
1546
|
+
visium_path
|
|
1547
|
+
Deprecated alias for ``spatial_path``, kept so existing notebooks keep
|
|
1548
|
+
working.
|
|
1549
|
+
|
|
1550
|
+
Returns
|
|
1551
|
+
-------
|
|
1552
|
+
Dict[str, SliceData]
|
|
1553
|
+
"""
|
|
1554
|
+
if visium_path is not None:
|
|
1555
|
+
if spatial_path is not None:
|
|
1556
|
+
raise TypeError(
|
|
1557
|
+
"step1_acquisition_and_anchoring: pass either spatial_path or the "
|
|
1558
|
+
"deprecated visium_path, not both."
|
|
1559
|
+
)
|
|
1560
|
+
warnings.warn(
|
|
1561
|
+
"visium_path is deprecated and will be removed in a future release; "
|
|
1562
|
+
"use spatial_path (the pipeline is no longer Visium-only).",
|
|
1563
|
+
DeprecationWarning, stacklevel=2,
|
|
1564
|
+
)
|
|
1565
|
+
spatial_path = visium_path
|
|
1566
|
+
if spatial_path is None:
|
|
1567
|
+
raise TypeError("step1_acquisition_and_anchoring: spatial_path is required.")
|
|
1568
|
+
|
|
1569
|
+
start_time = time.perf_counter()
|
|
1570
|
+
paths = _normalize_paths(spatial_path)
|
|
1571
|
+
platforms = _broadcast(platform, list(paths.keys()), default="auto")
|
|
1572
|
+
out: Dict[str, SliceData] = {}
|
|
1573
|
+
|
|
1574
|
+
for name, data_path in paths.items():
|
|
1575
|
+
plat = platforms[name]
|
|
1576
|
+
if plat == "auto":
|
|
1577
|
+
plat = detect_platform(data_path)
|
|
1578
|
+
if plat not in PLATFORM_PROFILES:
|
|
1579
|
+
raise ValueError(
|
|
1580
|
+
f"[{name}] unknown platform '{plat}'. "
|
|
1581
|
+
f"Known: {sorted(PLATFORM_PROFILES)}"
|
|
1582
|
+
)
|
|
1583
|
+
prof = PLATFORM_PROFILES[plat]
|
|
1584
|
+
print(f"\nStep 1 [{name}]: loading from {data_path}")
|
|
1585
|
+
print(
|
|
1586
|
+
f" Detected platform: {prof['label']}, "
|
|
1587
|
+
f"{prof['spot_diameter_um']:g} um capture units, "
|
|
1588
|
+
f"{prof['pitch_um']:g} um centre-to-centre"
|
|
1589
|
+
)
|
|
1590
|
+
|
|
1591
|
+
if plat != 'visium':
|
|
1592
|
+
out[name] = _load_generic_platform(name, data_path, plat)
|
|
1593
|
+
continue
|
|
1594
|
+
|
|
1595
|
+
out[name] = _load_visium_slice(name, data_path)
|
|
1596
|
+
|
|
1597
|
+
duration = time.perf_counter() - start_time
|
|
1598
|
+
print(f"\nStep 1 done in {duration:.2f}s {len(out)} slice(s) loaded.")
|
|
1599
|
+
return out
|
|
1600
|
+
|
|
1601
|
+
|
|
1602
|
+
def _load_visium_slice(name: str, data_path: str) -> SliceData:
|
|
1603
|
+
"""Load one 10x Visium slice from the SpaceRanger layout."""
|
|
1604
|
+
if True:
|
|
1605
|
+
spatial_dir = os.path.join(data_path, "spatial")
|
|
1606
|
+
h5_file = os.path.join(data_path, "filtered_feature_bc_matrix.h5")
|
|
1607
|
+
matrix_dir = os.path.join(data_path, "filtered_feature_bc_matrix")
|
|
1608
|
+
|
|
1609
|
+
if os.path.exists(h5_file):
|
|
1610
|
+
print(f" > h5 file: {h5_file}")
|
|
1611
|
+
with h5py.File(h5_file, 'r') as f:
|
|
1612
|
+
mat_group = f['matrix'] if 'matrix' in f else f
|
|
1613
|
+
data = mat_group['data'][:]
|
|
1614
|
+
indices = mat_group['indices'][:]
|
|
1615
|
+
indptr = mat_group['indptr'][:]
|
|
1616
|
+
shape = mat_group['shape'][:]
|
|
1617
|
+
counts = sp.csc_matrix((data, indices, indptr), shape=shape).T.tocsr()
|
|
1618
|
+
if 'features' in mat_group:
|
|
1619
|
+
feat_group = mat_group['features']
|
|
1620
|
+
gene_names = [x.decode('utf-8') for x in feat_group['name'][:]]
|
|
1621
|
+
else:
|
|
1622
|
+
gene_names = [x.decode('utf-8') for x in mat_group['genes'][:]]
|
|
1623
|
+
raw_barcodes = [x.decode('utf-8') for x in mat_group['barcodes'][:]]
|
|
1624
|
+
elif os.path.exists(matrix_dir):
|
|
1625
|
+
print(f" > matrix dir: {matrix_dir}")
|
|
1626
|
+
counts = scipy.io.mmread(os.path.join(matrix_dir, "matrix.mtx.gz")).T.tocsr()
|
|
1627
|
+
features = pd.read_csv(os.path.join(matrix_dir, "features.tsv.gz"), header=None, sep='\t')
|
|
1628
|
+
gene_names = features[1].values.tolist()
|
|
1629
|
+
raw_barcodes = pd.read_csv(
|
|
1630
|
+
os.path.join(matrix_dir, "barcodes.tsv.gz"), header=None, sep='\t'
|
|
1631
|
+
)[0].values.tolist()
|
|
1632
|
+
else:
|
|
1633
|
+
raise FileNotFoundError(
|
|
1634
|
+
f"[{name}] No filtered_feature_bc_matrix.h5 or filtered_feature_bc_matrix/ in {data_path}"
|
|
1635
|
+
)
|
|
1636
|
+
|
|
1637
|
+
# Spatial manifest
|
|
1638
|
+
pos_path = os.path.join(spatial_dir, "tissue_positions.csv")
|
|
1639
|
+
if not os.path.exists(pos_path):
|
|
1640
|
+
pos_path = os.path.join(spatial_dir, "tissue_positions_list.csv")
|
|
1641
|
+
has_header = 0
|
|
1642
|
+
if "list" in pos_path:
|
|
1643
|
+
has_header = None
|
|
1644
|
+
else:
|
|
1645
|
+
with open(pos_path, 'r') as f:
|
|
1646
|
+
first_line = f.readline()
|
|
1647
|
+
if "in_tissue" not in first_line and "barcode" not in first_line:
|
|
1648
|
+
has_header = None
|
|
1649
|
+
|
|
1650
|
+
spatial_df = pd.read_csv(pos_path, header=has_header)
|
|
1651
|
+
if len(spatial_df.columns) == 6:
|
|
1652
|
+
spatial_df.columns = ['barcode', 'in_tissue', 'array_row', 'array_col', 'pxl_row', 'pxl_col']
|
|
1653
|
+
else:
|
|
1654
|
+
spatial_df = spatial_df.rename(columns={
|
|
1655
|
+
'pxl_row_in_fullres': 'pxl_row',
|
|
1656
|
+
'pxl_col_in_fullres': 'pxl_col',
|
|
1657
|
+
})
|
|
1658
|
+
spatial_df = spatial_df.set_index('barcode')
|
|
1659
|
+
|
|
1660
|
+
with open(os.path.join(spatial_dir, "scalefactors_json.json"), 'r') as f:
|
|
1661
|
+
scale_factors = json.load(f)
|
|
1662
|
+
|
|
1663
|
+
# Barcode reconciliation
|
|
1664
|
+
if not any(b in spatial_df.index for b in raw_barcodes[:10]):
|
|
1665
|
+
print(" ! barcode mismatch; checking suffix conventions")
|
|
1666
|
+
if raw_barcodes[0].endswith("-1") and not spatial_df.index[0].endswith("-1"):
|
|
1667
|
+
raw_barcodes = [b.split('-')[0] for b in raw_barcodes]
|
|
1668
|
+
elif not raw_barcodes[0].endswith("-1") and spatial_df.index[0].endswith("-1"):
|
|
1669
|
+
raw_barcodes = [b + "-1" for b in raw_barcodes]
|
|
1670
|
+
|
|
1671
|
+
valid_barcodes = [b for b in raw_barcodes if b in spatial_df.index]
|
|
1672
|
+
valid_indices = [raw_barcodes.index(b) for b in valid_barcodes]
|
|
1673
|
+
if len(valid_barcodes) == 0:
|
|
1674
|
+
raise ValueError(f"[{name}] no common barcodes between matrix and spatial data")
|
|
1675
|
+
|
|
1676
|
+
counts = counts[valid_indices, :]
|
|
1677
|
+
coords = spatial_df.loc[valid_barcodes, ['pxl_row', 'pxl_col', 'array_row', 'array_col']].values
|
|
1678
|
+
total_counts_orig = np.array(counts.sum(axis=1)).flatten()
|
|
1679
|
+
|
|
1680
|
+
# QC plot per slice
|
|
1681
|
+
mean_umi = np.mean(total_counts_orig)
|
|
1682
|
+
median_umi = np.median(total_counts_orig)
|
|
1683
|
+
p1 = np.percentile(total_counts_orig, 1)
|
|
1684
|
+
print(f" median UMI/spot: {median_umi:.0f} mean: {mean_umi:.0f} 1st pct: {p1:.0f}")
|
|
1685
|
+
|
|
1686
|
+
plt.figure(figsize=(10, 6), dpi=100)
|
|
1687
|
+
bins = np.logspace(np.log10(max(1, total_counts_orig.min())), np.log10(total_counts_orig.max()), 50)
|
|
1688
|
+
plt.hist(total_counts_orig, bins=bins, color='#3498db', edgecolor='black', alpha=0.7)
|
|
1689
|
+
plt.axvline(median_umi, color='green', linestyle='--', label=f'Median: {int(median_umi)}')
|
|
1690
|
+
plt.axvline(p1, color='purple', linestyle=':', label=f'1st pct: {int(p1)}')
|
|
1691
|
+
plt.xscale('log')
|
|
1692
|
+
plt.title(f"[{name}] UMI per spot N={len(valid_barcodes)} spots")
|
|
1693
|
+
plt.xlabel("Total UMI (log)")
|
|
1694
|
+
plt.ylabel("Spots")
|
|
1695
|
+
plt.legend()
|
|
1696
|
+
plt.grid(True, which="both", ls="--", alpha=0.3)
|
|
1697
|
+
_show(f"step1_qc_umi_{name}")
|
|
1698
|
+
|
|
1699
|
+
return SliceData(
|
|
1700
|
+
name=name,
|
|
1701
|
+
spatial_path=spatial_dir,
|
|
1702
|
+
counts=counts,
|
|
1703
|
+
gene_names=gene_names,
|
|
1704
|
+
barcodes=valid_barcodes,
|
|
1705
|
+
coords=coords,
|
|
1706
|
+
total_umi=total_counts_orig,
|
|
1707
|
+
scale_factors=scale_factors,
|
|
1708
|
+
platform='visium',
|
|
1709
|
+
)
|
|
1710
|
+
|
|
1711
|
+
|
|
1712
|
+
def _load_generic_platform(name: str, data_path: str, plat: str) -> SliceData:
|
|
1713
|
+
"""Load a grid-indexed spot-based platform into the same SliceData shape.
|
|
1714
|
+
|
|
1715
|
+
Expects a single delimited counts table whose row (or column) labels encode
|
|
1716
|
+
the array coordinates as ``AxB`` / ``A_B``, which is the DBiT-seq layout.
|
|
1717
|
+
|
|
1718
|
+
Coordinates are converted to microns using the platform pitch and then used
|
|
1719
|
+
as the working pixel space at 1 px = 1 um, so ``spot_diameter_fullres`` and
|
|
1720
|
+
everything derived from it stay meaningful. No histology exists for these
|
|
1721
|
+
platforms, so ``spatial_path`` is left pointing at a directory that Step 9
|
|
1722
|
+
will populate with a black placeholder.
|
|
1723
|
+
"""
|
|
1724
|
+
prof = dict(PLATFORM_PROFILES[plat])
|
|
1725
|
+
|
|
1726
|
+
if os.path.isdir(data_path):
|
|
1727
|
+
tables = _find_count_tables(data_path)
|
|
1728
|
+
if not tables:
|
|
1729
|
+
raise FileNotFoundError(f"[{name}] no counts table found in {data_path}")
|
|
1730
|
+
table_path = tables[0]
|
|
1731
|
+
else:
|
|
1732
|
+
table_path = data_path
|
|
1733
|
+
|
|
1734
|
+
# DBiT-seq ships 10, 25 and 50 um variants; take the width from the sample
|
|
1735
|
+
# name when it carries one, and derive the pitch as twice the width.
|
|
1736
|
+
if plat == 'dbit_seq':
|
|
1737
|
+
px = _dbit_pixel_um(os.path.basename(table_path) or name, prof['spot_diameter_um'])
|
|
1738
|
+
prof['spot_diameter_um'], prof['pitch_um'] = px, 2.0 * px
|
|
1739
|
+
print(f" microchannel width {px:g} um -> pitch {2 * px:g} um")
|
|
1740
|
+
|
|
1741
|
+
sep = ',' if '.csv' in os.path.basename(table_path).lower() else '\t'
|
|
1742
|
+
print(f" > counts table: {table_path}")
|
|
1743
|
+
mat, row_labels, col_labels = _read_counts_table(table_path, sep)
|
|
1744
|
+
|
|
1745
|
+
# Orient to capture-units x genes: the capture-unit axis is the one whose
|
|
1746
|
+
# labels parse as coordinates, whichever way round the file was written.
|
|
1747
|
+
if _parse_coord_labels(row_labels[:50]) is None:
|
|
1748
|
+
mat, row_labels, col_labels = mat.T.tocsr(), col_labels, row_labels
|
|
1749
|
+
barcodes = row_labels
|
|
1750
|
+
xy = _parse_coord_labels(barcodes)
|
|
1751
|
+
if xy is None:
|
|
1752
|
+
raise ValueError(
|
|
1753
|
+
f"[{name}] labels of {table_path} are not coordinate-like "
|
|
1754
|
+
f"('AxB' or 'A_B') on either axis."
|
|
1755
|
+
)
|
|
1756
|
+
|
|
1757
|
+
gene_names = col_labels
|
|
1758
|
+
counts = mat.tocsr()
|
|
1759
|
+
|
|
1760
|
+
# Physical calibration. Grid indices scale by the known pitch directly;
|
|
1761
|
+
# otherwise recover the unit from the measured spacing (see PLATFORM_PROFILES).
|
|
1762
|
+
if prof['coord_units'] == 'grid':
|
|
1763
|
+
um_per_unit = float(prof['pitch_um'])
|
|
1764
|
+
array_idx = xy.copy()
|
|
1765
|
+
else:
|
|
1766
|
+
d_nn, _ = cKDTree(xy).query(xy, k=2)
|
|
1767
|
+
measured = float(np.median(d_nn[:, 1]))
|
|
1768
|
+
if not np.isfinite(measured) or measured <= 0:
|
|
1769
|
+
raise ValueError(f"[{name}] degenerate coordinates; cannot calibrate scale.")
|
|
1770
|
+
um_per_unit = float(prof['pitch_um']) / measured
|
|
1771
|
+
array_idx = xy / measured
|
|
1772
|
+
print(f" calibration: median nearest-neighbour spacing {measured:.2f} units "
|
|
1773
|
+
f"= {prof['pitch_um']:g} um -> {um_per_unit:.4f} um/unit")
|
|
1774
|
+
xy_um = xy * um_per_unit
|
|
1775
|
+
# Anchor at the origin so the placeholder canvas fits tightly.
|
|
1776
|
+
xy_um = xy_um - xy_um.min(axis=0) + float(prof['spot_diameter_um'])
|
|
1777
|
+
|
|
1778
|
+
coords = np.c_[xy_um[:, 0], xy_um[:, 1], array_idx[:, 0], array_idx[:, 1]]
|
|
1779
|
+
total_umi = np.asarray(counts.sum(axis=1)).flatten()
|
|
1780
|
+
|
|
1781
|
+
# Working pixel scale is 1 px = 1 um for these platforms, so the footprint in
|
|
1782
|
+
# "pixels" is the footprint in microns and cellchat_spatial_factors recovers
|
|
1783
|
+
# the ratio back as 1.0 from the same pitch rule.
|
|
1784
|
+
scale_factors = {
|
|
1785
|
+
'spot_diameter_fullres': float(prof['spot_diameter_um']),
|
|
1786
|
+
'fiducial_diameter_fullres': float(prof['spot_diameter_um']) * 1.75,
|
|
1787
|
+
'tissue_hires_scalef': 1.0,
|
|
1788
|
+
'tissue_lowres_scalef': 1.0,
|
|
1789
|
+
}
|
|
1790
|
+
|
|
1791
|
+
print(
|
|
1792
|
+
f" {counts.shape[0]} capture units x {counts.shape[1]} genes "
|
|
1793
|
+
f"median UMI/unit: {np.median(total_umi):.0f} mean: {total_umi.mean():.0f} "
|
|
1794
|
+
f"1st pct: {np.percentile(total_umi, 1):.0f} "
|
|
1795
|
+
f"extent {np.ptp(xy_um[:, 0]):.0f} x {np.ptp(xy_um[:, 1]):.0f} um"
|
|
1796
|
+
)
|
|
1797
|
+
# Only flag images that could plausibly be histology; instrument output
|
|
1798
|
+
# (bead frames, base-call images) is not tissue and naming it would mislead.
|
|
1799
|
+
folder = data_path if os.path.isdir(data_path) else os.path.dirname(data_path)
|
|
1800
|
+
sibling = [
|
|
1801
|
+
f for f in glob.glob(os.path.join(folder, '*'))
|
|
1802
|
+
if os.path.splitext(f)[1].lower() in ('.png', '.jpg', '.jpeg')
|
|
1803
|
+
and not re.search(r'worker|beadimage|calls|thumb', os.path.basename(f), re.I)
|
|
1804
|
+
]
|
|
1805
|
+
print(
|
|
1806
|
+
f" WARNING: no registered tissue image; a black placeholder will be written "
|
|
1807
|
+
f"at export so the object still loads in Seurat."
|
|
1808
|
+
)
|
|
1809
|
+
if sibling:
|
|
1810
|
+
print(
|
|
1811
|
+
f" (found {os.path.basename(sibling[0])} in the sample folder, but no "
|
|
1812
|
+
f"registration to the capture grid is provided, so it is not used.)"
|
|
1813
|
+
)
|
|
1814
|
+
|
|
1815
|
+
plt.figure(figsize=(10, 6), dpi=100)
|
|
1816
|
+
lo = max(1.0, float(total_umi.min()))
|
|
1817
|
+
bins = np.logspace(np.log10(lo), np.log10(max(lo + 1, float(total_umi.max()))), 50)
|
|
1818
|
+
plt.hist(total_umi, bins=bins, color='#3498db', edgecolor='black', alpha=0.7)
|
|
1819
|
+
plt.axvline(np.median(total_umi), color='green', linestyle='--',
|
|
1820
|
+
label=f'Median: {int(np.median(total_umi))}')
|
|
1821
|
+
plt.axvline(np.percentile(total_umi, 1), color='purple', linestyle=':',
|
|
1822
|
+
label=f'1st pct: {int(np.percentile(total_umi, 1))}')
|
|
1823
|
+
plt.xscale('log')
|
|
1824
|
+
plt.title(f"[{name}] UMI per capture unit N={counts.shape[0]} {prof['label']}")
|
|
1825
|
+
plt.xlabel("Total UMI (log)")
|
|
1826
|
+
plt.ylabel("Capture units")
|
|
1827
|
+
plt.legend()
|
|
1828
|
+
plt.grid(True, which="both", ls="--", alpha=0.3)
|
|
1829
|
+
_show(f"step1_qc_umi_{name}")
|
|
1830
|
+
|
|
1831
|
+
return SliceData(
|
|
1832
|
+
name=name,
|
|
1833
|
+
spatial_path=os.path.join(data_path if os.path.isdir(data_path)
|
|
1834
|
+
else os.path.dirname(data_path), "spatial"),
|
|
1835
|
+
counts=counts,
|
|
1836
|
+
gene_names=gene_names,
|
|
1837
|
+
barcodes=barcodes,
|
|
1838
|
+
coords=coords,
|
|
1839
|
+
total_umi=total_umi,
|
|
1840
|
+
scale_factors=scale_factors,
|
|
1841
|
+
platform=plat,
|
|
1842
|
+
)
|
|
1843
|
+
|
|
1844
|
+
|
|
1845
|
+
# STEP 2: density
|
|
1846
|
+
|
|
1847
|
+
@_profile_step
|
|
1848
|
+
def step2_estimate_cell_density(
|
|
1849
|
+
slices: Dict[str, SliceData],
|
|
1850
|
+
config_path: Optional[str] = None,
|
|
1851
|
+
species: str = "hs",
|
|
1852
|
+
min_umi=300,
|
|
1853
|
+
anchor_mean_factor=1.0,
|
|
1854
|
+
anchor_blend_alpha=0.6,
|
|
1855
|
+
low_slice_quality=False,
|
|
1856
|
+
colormap='hsv',
|
|
1857
|
+
) -> Dict[str, SliceData]:
|
|
1858
|
+
"""Estimate per-spot cell counts from Hybrid HK-UMI calibration.
|
|
1859
|
+
|
|
1860
|
+
Per-slice min_umi, anchor_mean_factor, and alpha accepted as scalar or dict.
|
|
1861
|
+
alpha=1.0 is pure UMI, alpha=0.0 is pure Housekeeping.
|
|
1862
|
+
"""
|
|
1863
|
+
start_time = time.perf_counter()
|
|
1864
|
+
profile = _load_config(config_path, species)
|
|
1865
|
+
hk_reference = profile['hk_profiles']
|
|
1866
|
+
engine_params = profile.get('engine_parameters', {})
|
|
1867
|
+
|
|
1868
|
+
if not hk_reference:
|
|
1869
|
+
raise ValueError(f"Species '{species}' has no hk_profiles in config.")
|
|
1870
|
+
|
|
1871
|
+
names = list(slices.keys())
|
|
1872
|
+
min_umi_d = _broadcast(min_umi, names)
|
|
1873
|
+
anchor_d = _broadcast(anchor_mean_factor, names)
|
|
1874
|
+
alpha_d = _broadcast(anchor_blend_alpha, names)
|
|
1875
|
+
low_q_d = _broadcast(low_slice_quality, names, default=False)
|
|
1876
|
+
|
|
1877
|
+
for name in names:
|
|
1878
|
+
sd = slices[name]
|
|
1879
|
+
slice_min_umi = min_umi_d[name]
|
|
1880
|
+
slice_anchor = anchor_d[name]
|
|
1881
|
+
slice_alpha = alpha_d[name]
|
|
1882
|
+
slice_low_q = low_q_d[name]
|
|
1883
|
+
|
|
1884
|
+
# anchor_blend_alpha is an internal, tuned constant and is deliberately
|
|
1885
|
+
# not reported: it is not a knob users should be reaching for.
|
|
1886
|
+
print(f"\nStep 2 [{name}]: Hybrid Calibration (factor={slice_anchor})")
|
|
1887
|
+
|
|
1888
|
+
common_hk = [g for g in hk_reference.keys() if g in sd.gene_names]
|
|
1889
|
+
hk_indices = [sd.gene_names.index(g) for g in common_hk]
|
|
1890
|
+
ref_values_log = np.array([hk_reference[g] for g in common_hk])
|
|
1891
|
+
|
|
1892
|
+
# 1. Housekeeping Signal (The Biological Anchor)
|
|
1893
|
+
safe_total = sd.total_umi.copy()
|
|
1894
|
+
safe_total[safe_total == 0] = 1
|
|
1895
|
+
hk_counts_raw = sd.counts[:, hk_indices].toarray()
|
|
1896
|
+
normalized_hk_log = np.log1p((hk_counts_raw / safe_total[:, np.newaxis]) * 10000)
|
|
1897
|
+
|
|
1898
|
+
spot_hk_log_means = np.mean(normalized_hk_log, axis=1)
|
|
1899
|
+
standard_anchor_log_mean = np.mean(ref_values_log)
|
|
1900
|
+
adjusted_standard_log_mean = standard_anchor_log_mean + np.log(slice_anchor)
|
|
1901
|
+
|
|
1902
|
+
spot_signal_linear = np.expm1(spot_hk_log_means)
|
|
1903
|
+
standard_signal_linear = np.expm1(adjusted_standard_log_mean)
|
|
1904
|
+
if standard_signal_linear < 0.001: standard_signal_linear = 0.001
|
|
1905
|
+
|
|
1906
|
+
hk_cells = spot_signal_linear / standard_signal_linear
|
|
1907
|
+
|
|
1908
|
+
# 2. Total UMI Signal (The Statistical Stabilizer)
|
|
1909
|
+
# We anchor the UMI-per-cell ratio to the valid HK-estimated spots
|
|
1910
|
+
valid_gate = sd.total_umi >= slice_min_umi
|
|
1911
|
+
if np.any(valid_gate) and np.sum(hk_cells[valid_gate]) > 0:
|
|
1912
|
+
global_umi_per_cell = np.sum(sd.total_umi[valid_gate]) / np.sum(hk_cells[valid_gate])
|
|
1913
|
+
else:
|
|
1914
|
+
global_umi_per_cell = np.mean(sd.total_umi) / 5.0 # Fallback
|
|
1915
|
+
|
|
1916
|
+
umi_cells = sd.total_umi / global_umi_per_cell
|
|
1917
|
+
|
|
1918
|
+
# 3. The Hybrid Blend (Weighted Geometric Mean)
|
|
1919
|
+
raw_n_cells = (umi_cells ** slice_alpha) * (hk_cells ** (1.0 - slice_alpha))
|
|
1920
|
+
|
|
1921
|
+
# 4. Discrete Mapping and Filtering
|
|
1922
|
+
n_cells = np.round(raw_n_cells).astype(int)
|
|
1923
|
+
is_low_quality = sd.total_umi < slice_min_umi
|
|
1924
|
+
n_cells[is_low_quality] = 0
|
|
1925
|
+
|
|
1926
|
+
if slice_low_q:
|
|
1927
|
+
floor_mask = (~is_low_quality) & (n_cells == 0)
|
|
1928
|
+
n_cells[floor_mask] = 1
|
|
1929
|
+
|
|
1930
|
+
# Calibration plot
|
|
1931
|
+
plt.figure(figsize=(8, 5), dpi=100)
|
|
1932
|
+
valid_signals = spot_signal_linear[sd.total_umi > slice_min_umi]
|
|
1933
|
+
max_x = max(np.percentile(valid_signals, 99) if len(valid_signals) > 0 else 10, standard_signal_linear * 2)
|
|
1934
|
+
bins = np.linspace(0, max_x, 50)
|
|
1935
|
+
plt.hist(valid_signals, bins=bins, color='purple', alpha=0.6, label='Spot HK signal')
|
|
1936
|
+
plt.axvline(standard_signal_linear, color='red', linewidth=2, linestyle='--',
|
|
1937
|
+
label=f'Standard ({standard_signal_linear:.2f})')
|
|
1938
|
+
plt.title(f"[{name}] Calibration check (factor={slice_anchor})")
|
|
1939
|
+
plt.xlabel("Geometric mean of HK normalized expression")
|
|
1940
|
+
plt.ylabel("Frequency")
|
|
1941
|
+
plt.legend()
|
|
1942
|
+
plt.grid(True, alpha=0.3)
|
|
1943
|
+
_show(f"step2_calibration_{name}")
|
|
1944
|
+
|
|
1945
|
+
n_zeros = int(np.sum(n_cells == 0))
|
|
1946
|
+
n_low = int(sum(is_low_quality))
|
|
1947
|
+
print(f" empty spots: {n_zeros}/{len(n_cells)} ({n_zeros/len(n_cells)*100:.1f}%) filtered by UMI: {n_low}")
|
|
1948
|
+
|
|
1949
|
+
# Density histogram
|
|
1950
|
+
plt.figure(figsize=(8, 5), dpi=100)
|
|
1951
|
+
max_val = int(np.max(n_cells)) if len(n_cells) > 0 else 0
|
|
1952
|
+
if max_val > 0:
|
|
1953
|
+
cell_bins = np.arange(1, max_val + 2) - 0.5
|
|
1954
|
+
plt.hist(n_cells[n_cells > 0], bins=cell_bins,
|
|
1955
|
+
color='#27ae60', edgecolor='white', alpha=0.9, label='Tissue')
|
|
1956
|
+
plt.bar(0, n_zeros, color='#95a5a6', edgecolor='white', width=0.8, label='Background')
|
|
1957
|
+
plt.title(f"[{name}] Cell density")
|
|
1958
|
+
plt.xlabel("Cells per spot")
|
|
1959
|
+
plt.xticks(np.arange(0, max(1, max_val) + 1, 1))
|
|
1960
|
+
plt.legend()
|
|
1961
|
+
plt.grid(axis='y', linestyle='--', alpha=0.3)
|
|
1962
|
+
_show(f"step2_density_hist_{name}")
|
|
1963
|
+
|
|
1964
|
+
# Spatial map
|
|
1965
|
+
y = sd.coords[:, 0]
|
|
1966
|
+
x = sd.coords[:, 1]
|
|
1967
|
+
plt.figure(figsize=(8, 8), dpi=100)
|
|
1968
|
+
bg_mask = n_cells == 0
|
|
1969
|
+
plt.scatter(x[bg_mask], y[bg_mask], c='grey', s=10, alpha=0.3)
|
|
1970
|
+
fg_mask = n_cells > 0
|
|
1971
|
+
if np.any(fg_mask):
|
|
1972
|
+
sc = plt.scatter(x[fg_mask], y[fg_mask], c=n_cells[fg_mask],
|
|
1973
|
+
cmap=colormap, s=15, linewidth=0)
|
|
1974
|
+
plt.colorbar(sc, label='Cells per spot', fraction=0.046, pad=0.04)
|
|
1975
|
+
plt.gca().invert_yaxis()
|
|
1976
|
+
plt.axis('off')
|
|
1977
|
+
plt.title(f"[{name}] spatial density")
|
|
1978
|
+
_show(f"step2_spatial_density_{name}")
|
|
1979
|
+
|
|
1980
|
+
print(f" total cells on slide: {int(sum(n_cells))}")
|
|
1981
|
+
|
|
1982
|
+
sd.n_cells = n_cells
|
|
1983
|
+
sd.engine_params = engine_params
|
|
1984
|
+
|
|
1985
|
+
duration = time.perf_counter() - start_time
|
|
1986
|
+
print(f"\nStep 2 done in {duration:.2f}s")
|
|
1987
|
+
return slices
|
|
1988
|
+
|
|
1989
|
+
|
|
1990
|
+
# STEP 3: ODG selection (joint, on intersected gene set)
|
|
1991
|
+
|
|
1992
|
+
@_profile_step
|
|
1993
|
+
def step3_feature_selection(
|
|
1994
|
+
slices: Dict[str, SliceData],
|
|
1995
|
+
config_path: Optional[str] = None,
|
|
1996
|
+
species: str = "hs",
|
|
1997
|
+
n_odg: int = 2000,
|
|
1998
|
+
od_poly_deg: int = 3,
|
|
1999
|
+
min_pct_spots: float = 0.05,
|
|
2000
|
+
max_pct_spots: float = 0.95,
|
|
2001
|
+
pct_filter_mode: str = "all",
|
|
2002
|
+
) -> OdgPack:
|
|
2003
|
+
"""Per-slice noise filter, intersect across slices, restrict the candidate
|
|
2004
|
+
gene pool by spot-presence fraction, then overdispersed-gene select on the
|
|
2005
|
+
candidates.
|
|
2006
|
+
|
|
2007
|
+
Pipeline inside this step:
|
|
2008
|
+
|
|
2009
|
+
1. Per-slice noise regex filter from the species profile.
|
|
2010
|
+
2. Intersect gene sets across slices.
|
|
2011
|
+
3. Restrict the candidate pool by spot-presence fraction (mimics
|
|
2012
|
+
STdeconvolve's restrictCorpus): per slice, the fraction of spots in
|
|
2013
|
+
which a gene has count > 0 must lie within
|
|
2014
|
+
``[min_pct_spots, max_pct_spots]``. With multiple slices, ``"any"``
|
|
2015
|
+
mode keeps a gene that passes in any one slice; ``"all"`` mode
|
|
2016
|
+
requires it to pass in every slice. The presence filter only restricts
|
|
2017
|
+
the pool considered for overdispersed selection, the per-slice
|
|
2018
|
+
``genes_clean`` and ``OdgPack.intersected_genes`` keep the full
|
|
2019
|
+
noise-filtered intersection, so Steps 4 / 6 / 7 / 9 still see all
|
|
2020
|
+
genes.
|
|
2021
|
+
4. Overdispersed-gene selection via mean-variance trend residual: the
|
|
2022
|
+
polynomial trend ``log10(var) ~ poly(log10(mean))`` is fit on
|
|
2023
|
+
candidates only (so rare and ubiquitous filtered-out genes don't pull
|
|
2024
|
+
the curve), and the top ``n_odg`` candidates by residual become the
|
|
2025
|
+
overdispersed gene set.
|
|
2026
|
+
|
|
2027
|
+
The OdgPack output carries the selected ``odg_names`` / ``odg_per_slice``
|
|
2028
|
+
/ ``odg_concat``; Steps 4-9 consume these directly as the overdispersed
|
|
2029
|
+
gene set.
|
|
2030
|
+
|
|
2031
|
+
Parameters
|
|
2032
|
+
----------
|
|
2033
|
+
n_odg : int
|
|
2034
|
+
Number of top overdispersed genes to retain.
|
|
2035
|
+
od_poly_deg : int
|
|
2036
|
+
Polynomial degree for the smoothed mean-variance trend (default 3).
|
|
2037
|
+
min_pct_spots, max_pct_spots : float
|
|
2038
|
+
Fraction-of-spots window for the presence filter. A gene passes if its
|
|
2039
|
+
spot-occupancy lies in [min_pct_spots, max_pct_spots] in (any | all)
|
|
2040
|
+
slice(s). Defaults 0.05 / 0.95.
|
|
2041
|
+
pct_filter_mode : str
|
|
2042
|
+
``"any"`` (default) or ``"all"``, see above.
|
|
2043
|
+
"""
|
|
2044
|
+
start_time = time.perf_counter()
|
|
2045
|
+
profile = _load_config(config_path, species)
|
|
2046
|
+
noise_pattern = profile.get('noise_regex', '')
|
|
2047
|
+
|
|
2048
|
+
print(f"Step 3: noise filter (species={species})")
|
|
2049
|
+
if noise_pattern:
|
|
2050
|
+
print(f" regex: {noise_pattern}")
|
|
2051
|
+
bio_re = re.compile(noise_pattern, re.IGNORECASE)
|
|
2052
|
+
else:
|
|
2053
|
+
bio_re = None
|
|
2054
|
+
print(" (no noise filter for this species)")
|
|
2055
|
+
|
|
2056
|
+
# Per-slice noise filter
|
|
2057
|
+
for name, sd in slices.items():
|
|
2058
|
+
if bio_re is None:
|
|
2059
|
+
keep_indices = list(range(len(sd.gene_names)))
|
|
2060
|
+
genes_kept = list(sd.gene_names)
|
|
2061
|
+
else:
|
|
2062
|
+
keep_indices = []
|
|
2063
|
+
genes_kept = []
|
|
2064
|
+
for i, g in enumerate(sd.gene_names):
|
|
2065
|
+
if not bio_re.match(g):
|
|
2066
|
+
keep_indices.append(i)
|
|
2067
|
+
genes_kept.append(g)
|
|
2068
|
+
sd.counts_clean = sd.counts[:, keep_indices]
|
|
2069
|
+
sd.genes_clean = genes_kept
|
|
2070
|
+
n_dropped = len(sd.gene_names) - len(genes_kept)
|
|
2071
|
+
print(f" [{name}] removed {n_dropped} noise genes; kept {len(genes_kept)}")
|
|
2072
|
+
|
|
2073
|
+
# Intersect gene sets
|
|
2074
|
+
if len(slices) == 1:
|
|
2075
|
+
only_name = next(iter(slices.keys()))
|
|
2076
|
+
intersected = list(slices[only_name].genes_clean)
|
|
2077
|
+
else:
|
|
2078
|
+
sets = [set(sd.genes_clean) for sd in slices.values()]
|
|
2079
|
+
intersected = sorted(set.intersection(*sets)) if len(sets) > 0 else []
|
|
2080
|
+
print(f"\nStep 3: gene set size = {len(intersected)}")
|
|
2081
|
+
if not intersected:
|
|
2082
|
+
raise ValueError("Empty intersection of slice gene sets!")
|
|
2083
|
+
|
|
2084
|
+
# Per-slice expression matrix on the intersected order
|
|
2085
|
+
intersected_set = set(intersected)
|
|
2086
|
+
inter_counts_per_slice: Dict[str, sp.csr_matrix] = {}
|
|
2087
|
+
for name, sd in slices.items():
|
|
2088
|
+
gene_to_idx = {g: i for i, g in enumerate(sd.genes_clean)}
|
|
2089
|
+
ord_idx = [gene_to_idx[g] for g in intersected]
|
|
2090
|
+
inter_counts_per_slice[name] = sd.counts_clean[:, ord_idx]
|
|
2091
|
+
|
|
2092
|
+
# Concatenate (rows = spots) for joint trend / overdispersed calculation.
|
|
2093
|
+
inter_concat = sp.vstack([inter_counts_per_slice[n] for n in slices.keys()]).tocsr()
|
|
2094
|
+
|
|
2095
|
+
# Presence filter (restrictCorpus). Per slice, compute the fraction of spots
|
|
2096
|
+
# in which each gene has count > 0; gene "passes" if the fraction is in
|
|
2097
|
+
# [min_pct_spots, max_pct_spots]. Combine across slices via the requested
|
|
2098
|
+
# mode. This filter only restricts the candidate pool for the overdispersed
|
|
2099
|
+
# selection, sd.genes_clean and OdgPack.intersected_genes remain untouched.
|
|
2100
|
+
if pct_filter_mode not in ("any", "all"):
|
|
2101
|
+
raise ValueError(
|
|
2102
|
+
f"pct_filter_mode must be 'any' or 'all', got {pct_filter_mode!r}"
|
|
2103
|
+
)
|
|
2104
|
+
|
|
2105
|
+
n_intersect = len(intersected)
|
|
2106
|
+
pass_matrix = np.zeros((len(slices), n_intersect), dtype=bool)
|
|
2107
|
+
print(f"Step 3: presence filter [{min_pct_spots:.2f}, {max_pct_spots:.2f}] "
|
|
2108
|
+
f"({pct_filter_mode} across slices)")
|
|
2109
|
+
for i, (name, _sd) in enumerate(slices.items()):
|
|
2110
|
+
X_s = inter_counts_per_slice[name]
|
|
2111
|
+
presence = np.asarray((X_s > 0).sum(axis=0)).flatten() / float(X_s.shape[0])
|
|
2112
|
+
passes = (presence >= min_pct_spots) & (presence <= max_pct_spots)
|
|
2113
|
+
pass_matrix[i] = passes
|
|
2114
|
+
print(f" [{name}] presence pass: {int(passes.sum())}/{n_intersect}")
|
|
2115
|
+
|
|
2116
|
+
if pct_filter_mode == "any":
|
|
2117
|
+
candidate_mask = pass_matrix.any(axis=0)
|
|
2118
|
+
else: # "all"
|
|
2119
|
+
candidate_mask = pass_matrix.all(axis=0)
|
|
2120
|
+
n_candidates = int(candidate_mask.sum())
|
|
2121
|
+
print(f" combined: {n_candidates}/{n_intersect} candidate genes after filter")
|
|
2122
|
+
|
|
2123
|
+
if n_candidates < max(10, od_poly_deg + 1):
|
|
2124
|
+
raise ValueError(
|
|
2125
|
+
f"Presence filter left only {n_candidates} candidate genes, "
|
|
2126
|
+
"loosen min_pct_spots / max_pct_spots or check the data."
|
|
2127
|
+
)
|
|
2128
|
+
|
|
2129
|
+
# Overdispersed gene selection: trend fit and ranking restricted to candidates.
|
|
2130
|
+
print(f"Step 3: selecting top {n_odg} overdispersed genes from {n_candidates} candidates...")
|
|
2131
|
+
n_odg_eff = min(n_odg, n_candidates)
|
|
2132
|
+
odg_indices_local, log_mean, log_var, fitted_log_var, residuals = _overdispersed_genes(
|
|
2133
|
+
inter_concat, n_top=n_odg_eff, poly_deg=od_poly_deg, fit_mask=candidate_mask
|
|
2134
|
+
)
|
|
2135
|
+
odg_names = [intersected[i] for i in odg_indices_local]
|
|
2136
|
+
|
|
2137
|
+
# Per-slice overdispersed-gene count matrices.
|
|
2138
|
+
odg_per_slice: Dict[str, sp.csr_matrix] = {
|
|
2139
|
+
name: inter_counts_per_slice[name][:, odg_indices_local] for name in slices.keys()
|
|
2140
|
+
}
|
|
2141
|
+
odg_concat = inter_concat[:, odg_indices_local]
|
|
2142
|
+
|
|
2143
|
+
# Diagnostic plot: residuals above the smoothed mean-variance trend, split
|
|
2144
|
+
# into three populations: filtered-out (failed presence filter), candidates
|
|
2145
|
+
# not selected, and selected overdispersed. x-axis clipped at the 5th
|
|
2146
|
+
# percentile so the +eps padded zero-expression tail doesn't blow out view.
|
|
2147
|
+
is_top = np.zeros(n_intersect, dtype=bool)
|
|
2148
|
+
is_top[odg_indices_local] = True
|
|
2149
|
+
cand_not_top = candidate_mask & (~is_top)
|
|
2150
|
+
filt_out = ~candidate_mask
|
|
2151
|
+
|
|
2152
|
+
plt.figure(figsize=(9, 6))
|
|
2153
|
+
plt.scatter(log_mean[filt_out], residuals[filt_out], s=2, color='lightgrey', alpha=0.35,
|
|
2154
|
+
label=f'Filtered out (presence): {int(filt_out.sum())}')
|
|
2155
|
+
plt.scatter(log_mean[cand_not_top], residuals[cand_not_top], s=2, color='steelblue', alpha=0.45,
|
|
2156
|
+
label=f'Candidates (not selected): {int(cand_not_top.sum())}')
|
|
2157
|
+
plt.scatter(log_mean[is_top], residuals[is_top], s=4, color='red',
|
|
2158
|
+
label=f'Overdispersed (selected): {n_odg_eff}')
|
|
2159
|
+
plt.axhline(0.0, color='black', linewidth=1.0, linestyle='--')
|
|
2160
|
+
|
|
2161
|
+
cand_mean = log_mean[candidate_mask]
|
|
2162
|
+
if cand_mean.size > 0:
|
|
2163
|
+
x_low = float(np.min(cand_mean)) - 0.2
|
|
2164
|
+
x_high = float(np.max(log_mean)) + 0.2
|
|
2165
|
+
plt.xlim(x_low, x_high)
|
|
2166
|
+
|
|
2167
|
+
cand_residuals = residuals[candidate_mask]
|
|
2168
|
+
if cand_residuals.size > 0:
|
|
2169
|
+
y_low = float(np.percentile(cand_residuals, 0.5))
|
|
2170
|
+
y_high = float(np.percentile(cand_residuals, 99.9))
|
|
2171
|
+
y_range = max(y_high - y_low, 1e-6)
|
|
2172
|
+
|
|
2173
|
+
pad = 0.1 * y_range
|
|
2174
|
+
plt.ylim(min(-0.5, y_low - pad), max(1.0, y_high + pad))
|
|
2175
|
+
|
|
2176
|
+
plt.xlabel('log10(mean expression)')
|
|
2177
|
+
plt.ylabel('Residual log10(variance): observed minus trend')
|
|
2178
|
+
plt.title(f'Step 3: overdispersion score (top {n_odg_eff} of {n_candidates} candidates; '
|
|
2179
|
+
f'{n_intersect} intersected genes)')
|
|
2180
|
+
plt.legend(loc='best')
|
|
2181
|
+
plt.grid(True, alpha=0.3)
|
|
2182
|
+
plt.tight_layout()
|
|
2183
|
+
_show("step3_overdispersion")
|
|
2184
|
+
|
|
2185
|
+
duration = time.perf_counter() - start_time
|
|
2186
|
+
print(f"Step 3 done in {duration:.2f}s joint overdispersed-gene matrix: {odg_concat.shape}")
|
|
2187
|
+
return OdgPack(
|
|
2188
|
+
intersected_genes=intersected,
|
|
2189
|
+
odg_names=odg_names,
|
|
2190
|
+
odg_per_slice=odg_per_slice,
|
|
2191
|
+
odg_concat=odg_concat,
|
|
2192
|
+
species=species,
|
|
2193
|
+
)
|
|
2194
|
+
|
|
2195
|
+
|
|
2196
|
+
# STEP 4: joint manifold
|
|
2197
|
+
|
|
2198
|
+
@_profile_step
|
|
2199
|
+
def step4_gene_manifold(
|
|
2200
|
+
slices: Dict[str, SliceData],
|
|
2201
|
+
odg_pack: OdgPack,
|
|
2202
|
+
config_path: Optional[str] = None,
|
|
2203
|
+
n_components: int = 30,
|
|
2204
|
+
) -> Manifold:
|
|
2205
|
+
"""ICA + UMAP on the gene matrix to produce a gene manifold.
|
|
2206
|
+
|
|
2207
|
+
Single slice: uses the slice's counts_clean directly in its native gene order
|
|
2208
|
+
(matches legacy codeconv Step 4 behavior bit-for-bit modulo upstream tie-breaking).
|
|
2209
|
+
Multi-slice: uses the alphabetical-intersected stacked matrix so all slices share
|
|
2210
|
+
a common gene index space.
|
|
2211
|
+
"""
|
|
2212
|
+
start_time = time.perf_counter()
|
|
2213
|
+
species = odg_pack.species
|
|
2214
|
+
profile = _load_config(config_path, species)
|
|
2215
|
+
qc_markers = profile.get('qc_markers', [])
|
|
2216
|
+
|
|
2217
|
+
n_slices = len(slices)
|
|
2218
|
+
if n_slices == 1:
|
|
2219
|
+
only_name = next(iter(slices.keys()))
|
|
2220
|
+
sd = slices[only_name]
|
|
2221
|
+
inter_concat = sd.counts_clean
|
|
2222
|
+
manifold_genes = list(sd.genes_clean)
|
|
2223
|
+
print(f"Step 4: building gene manifold over {inter_concat.shape[1]} genes...")
|
|
2224
|
+
else:
|
|
2225
|
+
inter_concat_per_slice = []
|
|
2226
|
+
for name, sd in slices.items():
|
|
2227
|
+
gene_to_idx = {g: i for i, g in enumerate(sd.genes_clean)}
|
|
2228
|
+
ord_idx = [gene_to_idx[g] for g in odg_pack.intersected_genes]
|
|
2229
|
+
inter_concat_per_slice.append(sd.counts_clean[:, ord_idx])
|
|
2230
|
+
inter_concat = sp.vstack(inter_concat_per_slice).tocsr()
|
|
2231
|
+
manifold_genes = list(odg_pack.intersected_genes)
|
|
2232
|
+
print(f"Step 4: building gene manifold over {inter_concat.shape[1]} intersected genes...")
|
|
2233
|
+
|
|
2234
|
+
# Transpose: now rows are genes, columns are spots-across-slices
|
|
2235
|
+
X_genes = inter_concat.T
|
|
2236
|
+
X_dense = np.log1p(X_genes.toarray())
|
|
2237
|
+
scaler = StandardScaler()
|
|
2238
|
+
X_scaled = scaler.fit_transform(X_dense)
|
|
2239
|
+
|
|
2240
|
+
print(f" FastICA n_components={n_components}")
|
|
2241
|
+
ica = FastICA(n_components=n_components, random_state=_SEED, max_iter=1000, tol=0.005)
|
|
2242
|
+
X_ica = ica.fit_transform(X_scaled)
|
|
2243
|
+
|
|
2244
|
+
print(" UMAP cosine projection")
|
|
2245
|
+
reducer = umap.UMAP(
|
|
2246
|
+
n_neighbors=30,
|
|
2247
|
+
min_dist=0.1,
|
|
2248
|
+
n_components=2,
|
|
2249
|
+
metric='cosine',
|
|
2250
|
+
random_state=_SEED,
|
|
2251
|
+
)
|
|
2252
|
+
embedding = reducer.fit_transform(X_ica)
|
|
2253
|
+
|
|
2254
|
+
# ODG positions within the manifold's gene index space
|
|
2255
|
+
inter_index = {g: i for i, g in enumerate(manifold_genes)}
|
|
2256
|
+
odg_indices_in_intersected = [inter_index[g] for g in odg_pack.odg_names]
|
|
2257
|
+
|
|
2258
|
+
# Visualization
|
|
2259
|
+
plt.figure(figsize=(8, 7), dpi=100)
|
|
2260
|
+
plt.scatter(embedding[:, 0], embedding[:, 1], s=2, c='lightgrey', alpha=0.4, label='Genes')
|
|
2261
|
+
|
|
2262
|
+
if qc_markers:
|
|
2263
|
+
colors = plt.cm.hsv(np.linspace(0, 1, len(qc_markers)))
|
|
2264
|
+
found = False
|
|
2265
|
+
for idx, gene in enumerate(qc_markers):
|
|
2266
|
+
if gene in inter_index:
|
|
2267
|
+
gi = inter_index[gene]
|
|
2268
|
+
plt.scatter(embedding[gi, 0], embedding[gi, 1],
|
|
2269
|
+
s=100, color=colors[idx], edgecolors='black', label=gene, zorder=10)
|
|
2270
|
+
plt.text(embedding[gi, 0] + 0.1, embedding[gi, 1] + 0.1, gene, fontsize=9)
|
|
2271
|
+
found = True
|
|
2272
|
+
if found:
|
|
2273
|
+
plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')
|
|
2274
|
+
else:
|
|
2275
|
+
print(" (no QC markers configured for this species; manifold shown un-annotated)")
|
|
2276
|
+
|
|
2277
|
+
plt.title("Step 4: joint gene co-expression manifold (ICA+UMAP)")
|
|
2278
|
+
plt.xlabel("UMAP 1")
|
|
2279
|
+
plt.ylabel("UMAP 2")
|
|
2280
|
+
plt.grid(True, alpha=0.15)
|
|
2281
|
+
plt.tight_layout()
|
|
2282
|
+
_show("step4_gene_manifold")
|
|
2283
|
+
|
|
2284
|
+
duration = time.perf_counter() - start_time
|
|
2285
|
+
print(f"Step 4 done in {duration:.2f}s")
|
|
2286
|
+
return Manifold(
|
|
2287
|
+
embedding=embedding,
|
|
2288
|
+
intersected_genes=manifold_genes,
|
|
2289
|
+
odg_indices_in_intersected=odg_indices_in_intersected,
|
|
2290
|
+
species=species,
|
|
2291
|
+
)
|
|
2292
|
+
|
|
2293
|
+
|
|
2294
|
+
# STEP 5: K sweep
|
|
2295
|
+
|
|
2296
|
+
@_profile_step
|
|
2297
|
+
def step5_ksweep(
|
|
2298
|
+
odg_pack: OdgPack,
|
|
2299
|
+
min_k: int = 3,
|
|
2300
|
+
max_k: int = 20,
|
|
2301
|
+
step: int = 1,
|
|
2302
|
+
subsample_frac: float = 1.0,
|
|
2303
|
+
doc_topic_prior: float = 0.1,
|
|
2304
|
+
topic_word_prior: float = 0.01,
|
|
2305
|
+
perc_rare_thresh: float = 0.05,
|
|
2306
|
+
alpha_mle_max_iter: int = 200,
|
|
2307
|
+
holdout_frac: float = 0.3,
|
|
2308
|
+
k_tol: float = 0.01,
|
|
2309
|
+
k_band: float = 0.05,
|
|
2310
|
+
k_rule: str = "percentile",
|
|
2311
|
+
n_topics: Optional[int] = None,
|
|
2312
|
+
) -> "KSweepResult":
|
|
2313
|
+
"""Run an LDA K-sweep and visualize held-out perplexity alongside rare-topic count.
|
|
2314
|
+
|
|
2315
|
+
For every K we fit LDA on a random 70/30 train/test split of the spots
|
|
2316
|
+
(split fixed across K so the K values are honestly comparable) and record:
|
|
2317
|
+
|
|
2318
|
+
1. Held-out perplexity, evaluated on the test split. Lower is a better
|
|
2319
|
+
out-of-sample fit. The train/test split prevents the same-data
|
|
2320
|
+
optimism the in-sample perplexity would have.
|
|
2321
|
+
2. Rare-topic count: number of topics whose mean proportion across the
|
|
2322
|
+
train spots is below ``perc_rare_thresh`` (default 0.05). Rare topics
|
|
2323
|
+
indicate the model is over-splitting into spurious cell types, K is
|
|
2324
|
+
"too high" once rare topics appear.
|
|
2325
|
+
3. Mean Dirichlet alpha, estimated post-hoc from the fitted theta via
|
|
2326
|
+
Minka's fixed-point iteration. Used internally by the recommended-K
|
|
2327
|
+
rule as a soft constraint (alpha < 1 means each spot stays dominated
|
|
2328
|
+
by a few topics). Not plotted because sklearn's fixed doc_topic_prior
|
|
2329
|
+
keeps the post-hoc estimate well below 1 across most K, so the metric
|
|
2330
|
+
is rarely informative in this setting; the values remain available on
|
|
2331
|
+
the returned object for inspection.
|
|
2332
|
+
|
|
2333
|
+
A K is recommended automatically, printed, and marked on the plot. The
|
|
2334
|
+
default rule is ``percentile`` at ``k_band=0.05``: take the largest K whose
|
|
2335
|
+
held-out perplexity is in the best 5% of the sweep by rank, the
|
|
2336
|
+
95th-percentile rule (see :func:`recommend_k`). This is the protocol's
|
|
2337
|
+
written guidance, highest K still in the low-perplexity regime, rare topics
|
|
2338
|
+
being desirable rather than penalized, expressed as a rule instead of a
|
|
2339
|
+
visual judgement call. Pass ``k_rule='relative_tolerance'`` for the
|
|
2340
|
+
sweep-extent-invariant variant, ``k_rule='largest_in_band'`` for the min-max
|
|
2341
|
+
band shaded on the plot, ``k_rule='conservative'`` for the module's original
|
|
2342
|
+
rule (lowest perplexity among K with alpha < 1 and no rare topics), or
|
|
2343
|
+
``k_rule='lowest_perplexity'`` for plain argmin.
|
|
2344
|
+
|
|
2345
|
+
The recommendation is a starting point, not a verdict: confirm it against
|
|
2346
|
+
the per-topic content in Step 6. To record a manual choice, pass it as
|
|
2347
|
+
``n_topics=``, it is echoed on the plot next to the recommendation and
|
|
2348
|
+
returned as ``chosen_k``, so ``ksweep.k`` is always the K the run proceeds
|
|
2349
|
+
with.
|
|
2350
|
+
|
|
2351
|
+
With multiple slices, also runs a joint sweep on the concatenated
|
|
2352
|
+
overdispersed-gene matrix and shows two cross-slice heatmaps: relative
|
|
2353
|
+
perplexity and rare-topic count. Single-slice runs skip the joint sweep and
|
|
2354
|
+
the heatmaps (both redundant).
|
|
2355
|
+
|
|
2356
|
+
Parameters
|
|
2357
|
+
----------
|
|
2358
|
+
perc_rare_thresh : float
|
|
2359
|
+
Threshold below which a topic's mean proportion across spots flags it
|
|
2360
|
+
as rare. Default 0.05.
|
|
2361
|
+
alpha_mle_max_iter : int
|
|
2362
|
+
Maximum iterations for the Minka fixed-point alpha estimator per K.
|
|
2363
|
+
holdout_frac : float
|
|
2364
|
+
Fraction of spots held out for perplexity evaluation. Default 0.3
|
|
2365
|
+
(70/30 train/test split). Pass 0.0 to evaluate perplexity on the full
|
|
2366
|
+
training matrix (the legacy in-sample behavior).
|
|
2367
|
+
k_band : float
|
|
2368
|
+
Band width as a fraction, for the default ``percentile`` rule (and for
|
|
2369
|
+
``largest_in_band`` and the shaded region on the plot). Default 0.05,
|
|
2370
|
+
the 95th-percentile rule. Raise it for a larger K.
|
|
2371
|
+
k_tol : float
|
|
2372
|
+
Relative tolerance on the best held-out perplexity, for the
|
|
2373
|
+
``relative_tolerance`` rule. Default 0.01 (within 1% of the best fit).
|
|
2374
|
+
k_rule : str
|
|
2375
|
+
``percentile`` (default) | ``relative_tolerance`` | ``largest_in_band`` |
|
|
2376
|
+
``lowest_perplexity`` | ``conservative``.
|
|
2377
|
+
n_topics : Optional[int]
|
|
2378
|
+
A manual K, if you have already decided. Does not change the sweep; it is
|
|
2379
|
+
recorded as ``chosen_k``, marked on the plot alongside the automatic
|
|
2380
|
+
recommendation, and returned by ``ksweep.k``.
|
|
2381
|
+
|
|
2382
|
+
Returns
|
|
2383
|
+
-------
|
|
2384
|
+
KSweepResult
|
|
2385
|
+
Dataclass with a compact ``__repr__`` so notebook auto-display stays one
|
|
2386
|
+
line. Fields: ``perplexity``, ``rare_topics``, ``alpha_mean``,
|
|
2387
|
+
``alpha_per_topic``, ``perc_rare_thresh``, ``recommended_k``, ``k_band``,
|
|
2388
|
+
``k_rule``, ``chosen_k``; ``.k`` gives the K to carry into Step 6 and
|
|
2389
|
+
``.summary()`` returns the whole sweep as a DataFrame.
|
|
2390
|
+
"""
|
|
2391
|
+
start_time = time.perf_counter()
|
|
2392
|
+
ks = list(range(min_k, max_k + 1, step))
|
|
2393
|
+
sweep_perp: Dict[str, List[float]] = {}
|
|
2394
|
+
sweep_rare: Dict[str, List[int]] = {}
|
|
2395
|
+
sweep_alpha_mean: Dict[str, List[float]] = {}
|
|
2396
|
+
sweep_alpha_full: Dict[str, List[np.ndarray]] = {}
|
|
2397
|
+
|
|
2398
|
+
n_slices = len(odg_pack.odg_per_slice)
|
|
2399
|
+
targets = dict(odg_pack.odg_per_slice)
|
|
2400
|
+
if n_slices > 1:
|
|
2401
|
+
targets['_joint'] = odg_pack.odg_concat
|
|
2402
|
+
|
|
2403
|
+
for label, X in targets.items():
|
|
2404
|
+
n_spots = X.shape[0]
|
|
2405
|
+
if subsample_frac < 1.0 and n_spots > 1000:
|
|
2406
|
+
n_sub = int(n_spots * subsample_frac)
|
|
2407
|
+
sub_rng = np.random.default_rng(_SEED)
|
|
2408
|
+
idx = sub_rng.choice(n_spots, n_sub, replace=False)
|
|
2409
|
+
X_use = X[idx, :]
|
|
2410
|
+
print(f"Step 5 [{label}]: {n_sub}/{n_spots} subsample")
|
|
2411
|
+
else:
|
|
2412
|
+
X_use = X
|
|
2413
|
+
print(f"Step 5 [{label}]: full {n_spots} spots")
|
|
2414
|
+
|
|
2415
|
+
# Train/test split. Fixed across K so the K values are honestly
|
|
2416
|
+
# comparable. holdout_frac=0.0 falls back to the legacy in-sample path.
|
|
2417
|
+
n_use = X_use.shape[0]
|
|
2418
|
+
if holdout_frac > 0.0 and n_use >= 10:
|
|
2419
|
+
split_rng = np.random.default_rng(_SEED + 7919)
|
|
2420
|
+
n_test = max(1, int(round(n_use * holdout_frac)))
|
|
2421
|
+
all_idx = np.arange(n_use)
|
|
2422
|
+
split_rng.shuffle(all_idx)
|
|
2423
|
+
test_idx = all_idx[:n_test]
|
|
2424
|
+
train_idx = all_idx[n_test:]
|
|
2425
|
+
X_train = X_use[train_idx, :]
|
|
2426
|
+
X_test = X_use[test_idx, :]
|
|
2427
|
+
print(f" train/test split: {len(train_idx)}/{len(test_idx)} ({(1 - holdout_frac):.2f}/{holdout_frac:.2f})")
|
|
2428
|
+
else:
|
|
2429
|
+
X_train = X_use
|
|
2430
|
+
X_test = X_use
|
|
2431
|
+
if holdout_frac > 0.0:
|
|
2432
|
+
print(f" too few spots ({n_use}) for a hold-out split; evaluating in-sample")
|
|
2433
|
+
|
|
2434
|
+
perps: List[float] = []
|
|
2435
|
+
rares: List[int] = []
|
|
2436
|
+
alpha_means: List[float] = []
|
|
2437
|
+
alpha_fulls: List[np.ndarray] = []
|
|
2438
|
+
for k in tqdm(ks, desc=f"LDA sweep [{label}]"):
|
|
2439
|
+
lda = LatentDirichletAllocation(
|
|
2440
|
+
n_components=k,
|
|
2441
|
+
learning_method='online',
|
|
2442
|
+
learning_offset=50.,
|
|
2443
|
+
max_iter=5,
|
|
2444
|
+
random_state=_SEED,
|
|
2445
|
+
n_jobs=-1,
|
|
2446
|
+
doc_topic_prior=doc_topic_prior,
|
|
2447
|
+
topic_word_prior=topic_word_prior,
|
|
2448
|
+
)
|
|
2449
|
+
# Fit on train; transform gives the train theta for the in-sample
|
|
2450
|
+
# metrics. Perplexity is evaluated on the held-out test split.
|
|
2451
|
+
theta_lda = lda.fit_transform(X_train)
|
|
2452
|
+
row_sum = theta_lda.sum(axis=1, keepdims=True)
|
|
2453
|
+
row_sum[row_sum == 0] = 1.0
|
|
2454
|
+
theta_props = theta_lda / row_sum
|
|
2455
|
+
|
|
2456
|
+
# Metric 1: held-out perplexity (on the test split).
|
|
2457
|
+
perps.append(lda.perplexity(X_test))
|
|
2458
|
+
# Metric 2: rare-topic count from train theta.
|
|
2459
|
+
mean_theta = theta_props.mean(axis=0)
|
|
2460
|
+
rares.append(int(np.sum(mean_theta < perc_rare_thresh)))
|
|
2461
|
+
# Metric 3: Minka MLE Dirichlet alpha from train theta.
|
|
2462
|
+
alpha_vec = _dirichlet_alpha_mle(theta_props, max_iter=alpha_mle_max_iter)
|
|
2463
|
+
alpha_fulls.append(alpha_vec)
|
|
2464
|
+
alpha_means.append(float(alpha_vec.mean()))
|
|
2465
|
+
|
|
2466
|
+
sweep_perp[label] = perps
|
|
2467
|
+
sweep_rare[label] = rares
|
|
2468
|
+
sweep_alpha_mean[label] = alpha_means
|
|
2469
|
+
sweep_alpha_full[label] = alpha_fulls
|
|
2470
|
+
|
|
2471
|
+
# Automatic K recommendation per target (see recommend_k for the rules).
|
|
2472
|
+
recommended_k: Dict[str, Dict[str, object]] = {}
|
|
2473
|
+
for label in sweep_perp.keys():
|
|
2474
|
+
recommended_k[label] = recommend_k(
|
|
2475
|
+
perplexity=dict(zip(ks, sweep_perp[label])),
|
|
2476
|
+
rule=k_rule,
|
|
2477
|
+
tol=k_tol,
|
|
2478
|
+
band=k_band,
|
|
2479
|
+
rare_topics=dict(zip(ks, sweep_rare[label])),
|
|
2480
|
+
alpha_mean=dict(zip(ks, sweep_alpha_mean[label])),
|
|
2481
|
+
)
|
|
2482
|
+
|
|
2483
|
+
# Report it. The previous version computed this silently, which made the
|
|
2484
|
+
# choice of K look like an unaided visual judgement call.
|
|
2485
|
+
headline_label = '_joint' if '_joint' in recommended_k else next(iter(recommended_k))
|
|
2486
|
+
headline = recommended_k[headline_label]
|
|
2487
|
+
|
|
2488
|
+
# Publish the choice as a module-level name as well as on the result object,
|
|
2489
|
+
# so Step 6 can read `codeconv.recommended_K` (or a bare `recommended_K`
|
|
2490
|
+
# after `from codeconv import *`) without threading the result through.
|
|
2491
|
+
global recommended_K
|
|
2492
|
+
recommended_K = int(n_topics) if n_topics is not None else int(headline['k'])
|
|
2493
|
+
|
|
2494
|
+
print(f"\nStep 5: recommended K = {headline['k']} [{headline['criterion']}]")
|
|
2495
|
+
if len(recommended_k) > 1:
|
|
2496
|
+
per_target = " ".join(f"{lbl}: K={rec['k']}" for lbl, rec in recommended_k.items())
|
|
2497
|
+
print(f" per target -> {per_target}")
|
|
2498
|
+
if headline['at_max_k']:
|
|
2499
|
+
print(
|
|
2500
|
+
f" WARNING: the recommendation sits at max_k={ks[-1]}, so perplexity had not "
|
|
2501
|
+
f"started to rise again, the sweep is too narrow to bound K from above. "
|
|
2502
|
+
f"Re-run with a larger max_k before trusting this value (Troubleshooting problem 7)."
|
|
2503
|
+
)
|
|
2504
|
+
if n_topics is not None:
|
|
2505
|
+
print(f" manual override: n_topics={int(n_topics)} will be used for Step 6.")
|
|
2506
|
+
|
|
2507
|
+
# Plot 1: dual-axis line plot, perplexity (solid o) + rare topics (right, dashed s)
|
|
2508
|
+
from matplotlib.ticker import MaxNLocator
|
|
2509
|
+
fig, ax1 = plt.subplots(figsize=(11, 6.3))
|
|
2510
|
+
color_cycle = plt.rcParams['axes.prop_cycle'].by_key()['color']
|
|
2511
|
+
label_color = {lbl: color_cycle[i % len(color_cycle)] for i, lbl in enumerate(sweep_perp.keys())}
|
|
2512
|
+
|
|
2513
|
+
for label, perps in sweep_perp.items():
|
|
2514
|
+
lw = 2.5 if label == '_joint' else 1.5
|
|
2515
|
+
ax1.plot(ks, perps, marker='o', markerfacecolor='white',
|
|
2516
|
+
color=label_color[label], linewidth=lw)
|
|
2517
|
+
ax1.set_xlabel("K topics")
|
|
2518
|
+
ax1.set_ylabel("Held-out perplexity (lower = better fit)")
|
|
2519
|
+
ax1.grid(True, alpha=0.3)
|
|
2520
|
+
|
|
2521
|
+
ax2 = ax1.twinx()
|
|
2522
|
+
for label, rares in sweep_rare.items():
|
|
2523
|
+
lw = 2.5 if label == '_joint' else 1.5
|
|
2524
|
+
ax2.plot(ks, rares, marker='s', linestyle='--',
|
|
2525
|
+
color=label_color[label], linewidth=lw, alpha=0.85)
|
|
2526
|
+
ax2.set_ylabel(f"# rare topics (mean θ < {perc_rare_thresh:.2f})")
|
|
2527
|
+
ax2.yaxis.set_major_locator(MaxNLocator(integer=True))
|
|
2528
|
+
|
|
2529
|
+
# Show the selection rule on the figure itself: the low-perplexity band as a
|
|
2530
|
+
# shaded region, the recommended K as a vertical line, a manual override (if
|
|
2531
|
+
# any) as a second line. Makes Figure 5 self-documenting.
|
|
2532
|
+
head_perp = np.array(sweep_perp[headline_label], dtype=float)
|
|
2533
|
+
p_lo, p_hi = float(head_perp.min()), float(head_perp.max())
|
|
2534
|
+
if p_hi > p_lo:
|
|
2535
|
+
if k_rule == "relative_tolerance":
|
|
2536
|
+
cutoff = p_lo * (1.0 + k_tol)
|
|
2537
|
+
band_label = f'within {k_tol:.1%} of best perplexity'
|
|
2538
|
+
elif k_rule == "percentile":
|
|
2539
|
+
cutoff = float(np.percentile(head_perp, k_band * 100.0))
|
|
2540
|
+
band_label = f'best {k_band:.0%} of sweep by rank'
|
|
2541
|
+
else:
|
|
2542
|
+
cutoff = p_lo + k_band * (p_hi - p_lo)
|
|
2543
|
+
band_label = f'low-perplexity band (≤{k_band:.0%} of range)'
|
|
2544
|
+
ax1.axhspan(p_lo, min(cutoff, p_hi), color='tab:green', alpha=0.08, zorder=0)
|
|
2545
|
+
ax1.axhline(cutoff, color='tab:green', linestyle=':', linewidth=1.2,
|
|
2546
|
+
label=band_label)
|
|
2547
|
+
ax1.axvline(headline['k'], color='tab:red', linestyle='--', linewidth=1.5,
|
|
2548
|
+
label=f"recommended K = {headline['k']}")
|
|
2549
|
+
if n_topics is not None and int(n_topics) != int(headline['k']):
|
|
2550
|
+
ax1.axvline(int(n_topics), color='tab:blue', linestyle='-.', linewidth=1.5,
|
|
2551
|
+
label=f"chosen K = {int(n_topics)}")
|
|
2552
|
+
|
|
2553
|
+
# Compact legend: one entry per slice (color) + a linestyle key for metric.
|
|
2554
|
+
slice_handles = [plt.Line2D([0], [0], color=label_color[lbl], linewidth=2, label=lbl)
|
|
2555
|
+
for lbl in sweep_perp.keys()]
|
|
2556
|
+
style_handles = [
|
|
2557
|
+
plt.Line2D([0], [0], color='black', linewidth=1.5,
|
|
2558
|
+
marker='o', markerfacecolor='white', label='Perplexity'),
|
|
2559
|
+
plt.Line2D([0], [0], color='black', linewidth=1.5,
|
|
2560
|
+
linestyle='--', marker='s', label='Rare topics'),
|
|
2561
|
+
]
|
|
2562
|
+
rule_handles, _ = ax1.get_legend_handles_labels()
|
|
2563
|
+
all_handles = slice_handles + style_handles + rule_handles
|
|
2564
|
+
# Centred above the axes so it never lands on the curves or the shaded band.
|
|
2565
|
+
ax1.legend(handles=all_handles, loc='lower center',
|
|
2566
|
+
bbox_to_anchor=(0.5, 1.03), ncol=min(4, len(all_handles)),
|
|
2567
|
+
fontsize=8, frameon=False)
|
|
2568
|
+
|
|
2569
|
+
title = (
|
|
2570
|
+
f"Step 5: K-sweep, perplexity and rare topics (recommended K = {headline['k']})"
|
|
2571
|
+
+ (", per slice + joint" if n_slices > 1 else "")
|
|
2572
|
+
)
|
|
2573
|
+
# Title above the legend, which itself sits above the axes.
|
|
2574
|
+
ax1.set_title(title, pad=52)
|
|
2575
|
+
plt.tight_layout()
|
|
2576
|
+
_show("step5_ksweep_perplexity_rare")
|
|
2577
|
+
|
|
2578
|
+
# Cross-slice heatmaps (only meaningful with 2+ slices).
|
|
2579
|
+
slice_labels = [k for k in sweep_perp.keys() if k != '_joint']
|
|
2580
|
+
if len(slice_labels) > 1:
|
|
2581
|
+
# Plot 2a: relative perplexity heatmap
|
|
2582
|
+
mat = np.zeros((len(slice_labels), len(ks)))
|
|
2583
|
+
for i, lbl in enumerate(slice_labels):
|
|
2584
|
+
row = np.array(sweep_perp[lbl], dtype=float)
|
|
2585
|
+
row_range = row.max() - row.min()
|
|
2586
|
+
mat[i] = (row - row.min()) / row_range if row_range > 0 else 0.0
|
|
2587
|
+
plt.figure(figsize=(max(8, len(ks) * 0.5), 0.6 * len(slice_labels) + 2))
|
|
2588
|
+
sns.heatmap(mat, xticklabels=ks, yticklabels=slice_labels,
|
|
2589
|
+
cmap='viridis_r', annot=False, cbar_kws={'label': 'relative perplexity'})
|
|
2590
|
+
raw_mat = np.array([sweep_perp[lbl] for lbl in slice_labels])
|
|
2591
|
+
median_per_k = np.median(raw_mat, axis=0)
|
|
2592
|
+
best_k_idx = int(np.argmin(median_per_k))
|
|
2593
|
+
plt.axvline(best_k_idx + 0.5, color='red', linestyle='--', linewidth=1.5)
|
|
2594
|
+
plt.title(f"Step 5: relative perplexity by K (red line = median-best K = {ks[best_k_idx]})")
|
|
2595
|
+
plt.xlabel("K")
|
|
2596
|
+
plt.ylabel("Slice")
|
|
2597
|
+
plt.tight_layout()
|
|
2598
|
+
_show("step5_relative_perplexity_heatmap")
|
|
2599
|
+
|
|
2600
|
+
# Plot 2b: rare-topic count heatmap (raw counts, annotated)
|
|
2601
|
+
rare_mat = np.array([sweep_rare[lbl] for lbl in slice_labels], dtype=int)
|
|
2602
|
+
plt.figure(figsize=(max(8, len(ks) * 0.5), 0.6 * len(slice_labels) + 2))
|
|
2603
|
+
sns.heatmap(rare_mat, xticklabels=ks, yticklabels=slice_labels,
|
|
2604
|
+
cmap='Reds', annot=True, fmt='d',
|
|
2605
|
+
cbar_kws={'label': f'# rare topics (mean θ < {perc_rare_thresh:.2f})'})
|
|
2606
|
+
median_rare_per_k = np.median(rare_mat, axis=0)
|
|
2607
|
+
zero_rare = np.where(median_rare_per_k == 0)[0]
|
|
2608
|
+
if len(zero_rare) > 0:
|
|
2609
|
+
last_clean_k_idx = int(zero_rare.max())
|
|
2610
|
+
plt.axvline(last_clean_k_idx + 0.5, color='blue', linestyle='--', linewidth=1.5)
|
|
2611
|
+
rare_guidance = f"largest K with median rare=0 is K={ks[last_clean_k_idx]}"
|
|
2612
|
+
else:
|
|
2613
|
+
rare_guidance = "all K have rare topics, consider lowering K"
|
|
2614
|
+
plt.title(f"Step 5: rare topics by K ({rare_guidance})")
|
|
2615
|
+
plt.xlabel("K")
|
|
2616
|
+
plt.ylabel("Slice")
|
|
2617
|
+
plt.tight_layout()
|
|
2618
|
+
_show("step5_rare_topics_heatmap")
|
|
2619
|
+
|
|
2620
|
+
duration = time.perf_counter() - start_time
|
|
2621
|
+
print(f"\nStep 5 done in {duration:.2f}s")
|
|
2622
|
+
return KSweepResult(
|
|
2623
|
+
perplexity={label: dict(zip(ks, perps)) for label, perps in sweep_perp.items()},
|
|
2624
|
+
rare_topics={label: dict(zip(ks, rares)) for label, rares in sweep_rare.items()},
|
|
2625
|
+
alpha_mean={label: dict(zip(ks, alphas)) for label, alphas in sweep_alpha_mean.items()},
|
|
2626
|
+
alpha_per_topic={label: dict(zip(ks, sweep_alpha_full[label])) for label in sweep_alpha_full},
|
|
2627
|
+
perc_rare_thresh=perc_rare_thresh,
|
|
2628
|
+
recommended_k=recommended_k,
|
|
2629
|
+
k_tol=k_tol,
|
|
2630
|
+
k_band=k_band,
|
|
2631
|
+
k_rule=k_rule,
|
|
2632
|
+
chosen_k=int(n_topics) if n_topics is not None else None,
|
|
2633
|
+
)
|
|
2634
|
+
|
|
2635
|
+
|
|
2636
|
+
# STEP 6: per-slice LDA -> Hungarian alignment -> consensus beta -> per-slice theta refit + projection
|
|
2637
|
+
|
|
2638
|
+
@_profile_step
|
|
2639
|
+
def step6_final_deconvolution(
|
|
2640
|
+
slices: Dict[str, SliceData],
|
|
2641
|
+
odg_pack: OdgPack,
|
|
2642
|
+
manifold: Manifold,
|
|
2643
|
+
n_topics: int,
|
|
2644
|
+
n_iters: int = 30,
|
|
2645
|
+
k_neighbors: int = 3,
|
|
2646
|
+
min_sim: float = 0.05,
|
|
2647
|
+
doc_topic_prior: float = 0.1,
|
|
2648
|
+
topic_word_prior: float = 0.01,
|
|
2649
|
+
e_step_iters: int = 50,
|
|
2650
|
+
n_seeds: int = 1,
|
|
2651
|
+
) -> Model:
|
|
2652
|
+
"""LDA fit per slice with optional multi-seed consensus to stabilize topics.
|
|
2653
|
+
|
|
2654
|
+
Per slice, fit LDA with ``n_seeds`` independent random initializations
|
|
2655
|
+
(seeds ``_SEED`` ... ``_SEED + n_seeds - 1``), Hungarian-align all replicate
|
|
2656
|
+
betas (and thetas) to the seed-0 ordering, and take their mean as that
|
|
2657
|
+
slice's beta. Single LDA fits are highly sensitive to initialization, so
|
|
2658
|
+
a 5-10 seed consensus gives substantially more reproducible topic content
|
|
2659
|
+
at the cost of an N-times longer Step 6.
|
|
2660
|
+
|
|
2661
|
+
A per-slice topic stability score, the mean pairwise cosine similarity of
|
|
2662
|
+
aligned topic rows across seeds, is computed and stored on the returned
|
|
2663
|
+
Model. 1.0 means seeds are identical; values closer to 0 mean topics drift
|
|
2664
|
+
across seeds and the K may be too high (or the data underdetermines that
|
|
2665
|
+
many topics).
|
|
2666
|
+
|
|
2667
|
+
With multiple slices, the per-slice consensus betas are then Hungarian-
|
|
2668
|
+
aligned across slices and meaned to form ``beta_consensus``, and per-slice
|
|
2669
|
+
theta is refit against the frozen consensus beta via the variational E-step.
|
|
2670
|
+
With a single slice, the per-slice consensus beta IS the model.
|
|
2671
|
+
Per-slice manifold-guided projection extends consensus topics to each
|
|
2672
|
+
slice's full ``genes_clean``.
|
|
2673
|
+
|
|
2674
|
+
Parameters
|
|
2675
|
+
----------
|
|
2676
|
+
n_seeds : int
|
|
2677
|
+
Number of LDA seeds per slice to consensus-average. Default 1 keeps the
|
|
2678
|
+
legacy single-seed behavior. 5-10 is recommended for stable topics.
|
|
2679
|
+
"""
|
|
2680
|
+
start_time = time.perf_counter()
|
|
2681
|
+
n_slices = len(slices)
|
|
2682
|
+
n_odg = len(odg_pack.odg_names)
|
|
2683
|
+
n_seeds = max(1, int(n_seeds))
|
|
2684
|
+
print(f"Step 6: LDA fit, K={n_topics} ({n_slices} slice{'s' if n_slices > 1 else ''}, n_seeds={n_seeds})")
|
|
2685
|
+
|
|
2686
|
+
per_slice_betas_unaligned: Dict[str, np.ndarray] = {}
|
|
2687
|
+
per_slice_thetas_lda: Dict[str, np.ndarray] = {}
|
|
2688
|
+
per_slice_stability: Dict[str, float] = {}
|
|
2689
|
+
for name in slices.keys():
|
|
2690
|
+
X = odg_pack.odg_per_slice[name]
|
|
2691
|
+
if n_seeds == 1:
|
|
2692
|
+
print(f" [{name}] LDA on {X.shape[0]} spots (single seed)")
|
|
2693
|
+
lda = LatentDirichletAllocation(
|
|
2694
|
+
n_components=n_topics,
|
|
2695
|
+
learning_method='batch',
|
|
2696
|
+
max_iter=n_iters,
|
|
2697
|
+
random_state=_SEED,
|
|
2698
|
+
n_jobs=-1,
|
|
2699
|
+
verbose=0,
|
|
2700
|
+
doc_topic_prior=doc_topic_prior,
|
|
2701
|
+
topic_word_prior=topic_word_prior,
|
|
2702
|
+
)
|
|
2703
|
+
theta_lda = lda.fit_transform(X)
|
|
2704
|
+
beta = lda.components_ / lda.components_.sum(axis=1, keepdims=True)
|
|
2705
|
+
per_slice_betas_unaligned[name] = beta
|
|
2706
|
+
per_slice_thetas_lda[name] = theta_lda
|
|
2707
|
+
per_slice_stability[name] = 1.0
|
|
2708
|
+
else:
|
|
2709
|
+
print(f" [{name}] LDA on {X.shape[0]} spots ({n_seeds}-seed consensus)")
|
|
2710
|
+
seed_betas: List[np.ndarray] = []
|
|
2711
|
+
seed_thetas: List[np.ndarray] = []
|
|
2712
|
+
for i in range(n_seeds):
|
|
2713
|
+
lda = LatentDirichletAllocation(
|
|
2714
|
+
n_components=n_topics,
|
|
2715
|
+
learning_method='batch',
|
|
2716
|
+
max_iter=n_iters,
|
|
2717
|
+
random_state=_SEED + i,
|
|
2718
|
+
n_jobs=-1,
|
|
2719
|
+
verbose=0,
|
|
2720
|
+
doc_topic_prior=doc_topic_prior,
|
|
2721
|
+
topic_word_prior=topic_word_prior,
|
|
2722
|
+
)
|
|
2723
|
+
t = lda.fit_transform(X)
|
|
2724
|
+
b = lda.components_ / lda.components_.sum(axis=1, keepdims=True)
|
|
2725
|
+
seed_betas.append(b)
|
|
2726
|
+
seed_thetas.append(t)
|
|
2727
|
+
|
|
2728
|
+
# Hungarian-align each seed to seed 0 on cosine of beta rows. Apply
|
|
2729
|
+
# the same column permutation to the corresponding theta.
|
|
2730
|
+
anchor = seed_betas[0]
|
|
2731
|
+
anchor_norms = np.linalg.norm(anchor, axis=1, keepdims=True)
|
|
2732
|
+
anchor_norms[anchor_norms == 0] = 1.0
|
|
2733
|
+
anchor_norm = anchor / anchor_norms
|
|
2734
|
+
aligned_b: List[np.ndarray] = [anchor]
|
|
2735
|
+
aligned_t: List[np.ndarray] = [seed_thetas[0]]
|
|
2736
|
+
for i in range(1, n_seeds):
|
|
2737
|
+
b = seed_betas[i]
|
|
2738
|
+
b_norms = np.linalg.norm(b, axis=1, keepdims=True)
|
|
2739
|
+
b_norms[b_norms == 0] = 1.0
|
|
2740
|
+
b_norm = b / b_norms
|
|
2741
|
+
sim = anchor_norm @ b_norm.T
|
|
2742
|
+
row_ind, col_ind = linear_sum_assignment(-sim)
|
|
2743
|
+
new_b = np.zeros_like(b)
|
|
2744
|
+
new_b[row_ind] = b[col_ind]
|
|
2745
|
+
new_t = seed_thetas[i][:, col_ind]
|
|
2746
|
+
aligned_b.append(new_b)
|
|
2747
|
+
aligned_t.append(new_t)
|
|
2748
|
+
|
|
2749
|
+
beta_mean = np.mean(np.stack(aligned_b, axis=0), axis=0)
|
|
2750
|
+
beta_mean = beta_mean / beta_mean.sum(axis=1, keepdims=True)
|
|
2751
|
+
theta_mean = np.mean(np.stack(aligned_t, axis=0), axis=0)
|
|
2752
|
+
t_sum = theta_mean.sum(axis=1, keepdims=True)
|
|
2753
|
+
t_sum[t_sum == 0] = 1.0
|
|
2754
|
+
theta_mean = theta_mean / t_sum
|
|
2755
|
+
|
|
2756
|
+
stability = _topic_stability(aligned_b)
|
|
2757
|
+
per_slice_betas_unaligned[name] = beta_mean
|
|
2758
|
+
per_slice_thetas_lda[name] = theta_mean
|
|
2759
|
+
per_slice_stability[name] = stability
|
|
2760
|
+
print(f" topic stability across {n_seeds} seeds: {stability:.3f} (1.0 = identical)")
|
|
2761
|
+
|
|
2762
|
+
if n_slices > 1:
|
|
2763
|
+
# Multi-slice path: align topics, build consensus beta, refit theta with frozen beta
|
|
2764
|
+
print("Step 6: Hungarian alignment of topics across slices")
|
|
2765
|
+
aligned_betas = _align_topics(per_slice_betas_unaligned)
|
|
2766
|
+
|
|
2767
|
+
print("Step 6: building consensus beta (mean across aligned slices)")
|
|
2768
|
+
beta_stack = np.stack([aligned_betas[n] for n in slices.keys()], axis=0)
|
|
2769
|
+
beta_consensus = beta_stack.mean(axis=0)
|
|
2770
|
+
beta_consensus = beta_consensus / beta_consensus.sum(axis=1, keepdims=True)
|
|
2771
|
+
|
|
2772
|
+
print(f"Step 6: per-slice theta refit (variational E-step, max_iter={e_step_iters})")
|
|
2773
|
+
for name, sd in slices.items():
|
|
2774
|
+
X = odg_pack.odg_per_slice[name]
|
|
2775
|
+
theta = _variational_e_step(X, beta_consensus, alpha=doc_topic_prior, max_iter=e_step_iters)
|
|
2776
|
+
sd.theta = theta
|
|
2777
|
+
print(f" [{name}] theta {theta.shape}")
|
|
2778
|
+
# Multi-slice projection: per-slice with cosine fallback for slice-specific genes
|
|
2779
|
+
print("Step 6: projecting topics to full genes_clean per slice")
|
|
2780
|
+
inter_index = {g: i for i, g in enumerate(manifold.intersected_genes)}
|
|
2781
|
+
odg_idx_in_intersected = manifold.odg_indices_in_intersected
|
|
2782
|
+
odg_coords_manifold = manifold.embedding[odg_idx_in_intersected]
|
|
2783
|
+
|
|
2784
|
+
per_slice_betas_full: Dict[str, np.ndarray] = {}
|
|
2785
|
+
for name, sd in slices.items():
|
|
2786
|
+
genes_s = sd.genes_clean
|
|
2787
|
+
n_genes_s = len(genes_s)
|
|
2788
|
+
beta_full = np.zeros((n_topics, n_genes_s))
|
|
2789
|
+
|
|
2790
|
+
slice_gene_to_idx = {g: i for i, g in enumerate(genes_s)}
|
|
2791
|
+
odg_idx_in_slice = [slice_gene_to_idx[g] for g in odg_pack.odg_names]
|
|
2792
|
+
|
|
2793
|
+
norm_all = normalize(sd.counts_clean.T, axis=1)
|
|
2794
|
+
norm_odg = normalize(sd.counts_clean[:, odg_idx_in_slice].T, axis=1)
|
|
2795
|
+
sim_full = norm_all @ norm_odg.T
|
|
2796
|
+
|
|
2797
|
+
for i, g in enumerate(genes_s):
|
|
2798
|
+
sim_row = sim_full[i].toarray().flatten() if sp.issparse(sim_full) else sim_full[i]
|
|
2799
|
+
if g in inter_index:
|
|
2800
|
+
gi = inter_index[g]
|
|
2801
|
+
g_pos = manifold.embedding[gi]
|
|
2802
|
+
dists = np.linalg.norm(odg_coords_manifold - g_pos, axis=1)
|
|
2803
|
+
neighbors = np.argsort(dists)[:k_neighbors]
|
|
2804
|
+
else:
|
|
2805
|
+
neighbors = np.argsort(-sim_row)[:k_neighbors]
|
|
2806
|
+
|
|
2807
|
+
weights = sim_row[neighbors]
|
|
2808
|
+
if np.max(weights) < min_sim:
|
|
2809
|
+
continue
|
|
2810
|
+
proj = beta_consensus[:, neighbors] @ weights
|
|
2811
|
+
if proj.sum() > 0:
|
|
2812
|
+
beta_full[:, i] = proj / proj.sum()
|
|
2813
|
+
|
|
2814
|
+
sd.beta_final = beta_full
|
|
2815
|
+
per_slice_betas_full[name] = beta_full
|
|
2816
|
+
print(f" [{name}] beta_final {beta_full.shape}")
|
|
2817
|
+
|
|
2818
|
+
# Multi-slice QC: rank genes per topic by log2 fold change of beta vs
|
|
2819
|
+
# the mean beta of the other topics. This highlights topic-specific
|
|
2820
|
+
# markers rather than genes that are simply highly expressed everywhere.
|
|
2821
|
+
log2fc_consensus = _topic_log2fc(beta_consensus)
|
|
2822
|
+
topic_dict = {}
|
|
2823
|
+
for k in range(n_topics):
|
|
2824
|
+
top_idx = log2fc_consensus[k].argsort()[::-1][:15]
|
|
2825
|
+
topic_dict[f"Topic_{k}"] = [odg_pack.odg_names[i] for i in top_idx]
|
|
2826
|
+
qc_df = pd.DataFrame(topic_dict)
|
|
2827
|
+
else:
|
|
2828
|
+
# Single-slice path mirrors the legacy codeconv Step 6: LDA fit_transform
|
|
2829
|
+
# gives theta and beta_odg, then a single cdist-based projection extends
|
|
2830
|
+
# beta to the full genes_clean. QC table is built from the projected beta_final.
|
|
2831
|
+
only_name = next(iter(slices.keys()))
|
|
2832
|
+
sd = slices[only_name]
|
|
2833
|
+
sd.theta = per_slice_thetas_lda[only_name]
|
|
2834
|
+
beta_odg = per_slice_betas_unaligned[only_name]
|
|
2835
|
+
beta_consensus = beta_odg
|
|
2836
|
+
print(f" [{only_name}] theta {sd.theta.shape}")
|
|
2837
|
+
|
|
2838
|
+
print("Step 6: projecting latent topics using UMAP manifold topology")
|
|
2839
|
+
genes_s = sd.genes_clean
|
|
2840
|
+
n_genes_s = len(genes_s)
|
|
2841
|
+
|
|
2842
|
+
# Manifold positions: for a single slice, every gene in genes_clean is in
|
|
2843
|
+
# the intersected set (intersection of one set is itself), so all genes have
|
|
2844
|
+
# manifold coordinates.
|
|
2845
|
+
inter_index = {g: i for i, g in enumerate(manifold.intersected_genes)}
|
|
2846
|
+
odg_coords_manifold = manifold.embedding[manifold.odg_indices_in_intersected]
|
|
2847
|
+
all_coords = manifold.embedding[[inter_index[g] for g in genes_s]]
|
|
2848
|
+
dist_matrix = cdist(all_coords, odg_coords_manifold, metric='euclidean')
|
|
2849
|
+
|
|
2850
|
+
# Cosine similarity in expression space
|
|
2851
|
+
slice_gene_to_idx = {g: i for i, g in enumerate(genes_s)}
|
|
2852
|
+
odg_idx_in_slice = [slice_gene_to_idx[g] for g in odg_pack.odg_names]
|
|
2853
|
+
norm_all = normalize(sd.counts_clean.T, axis=1)
|
|
2854
|
+
norm_odg = normalize(sd.counts_clean[:, odg_idx_in_slice].T, axis=1)
|
|
2855
|
+
similarity = norm_all @ norm_odg.T
|
|
2856
|
+
|
|
2857
|
+
beta_final = np.zeros((n_topics, n_genes_s))
|
|
2858
|
+
for i in range(n_genes_s):
|
|
2859
|
+
umap_neighbors_idx = np.argsort(dist_matrix[i])[:k_neighbors]
|
|
2860
|
+
sim_row = similarity[i].toarray().flatten() if sp.issparse(similarity) else similarity[i]
|
|
2861
|
+
weights = sim_row[umap_neighbors_idx]
|
|
2862
|
+
if np.max(weights) < min_sim:
|
|
2863
|
+
continue
|
|
2864
|
+
proj = beta_odg[:, umap_neighbors_idx] @ weights
|
|
2865
|
+
if proj.sum() > 0:
|
|
2866
|
+
beta_final[:, i] = proj / proj.sum()
|
|
2867
|
+
|
|
2868
|
+
sd.beta_final = beta_final
|
|
2869
|
+
per_slice_betas_full = {only_name: beta_final}
|
|
2870
|
+
print(f" [{only_name}] beta_final {beta_final.shape}")
|
|
2871
|
+
|
|
2872
|
+
# Single-slice QC: rank genes per topic by log2 fold change of beta_final
|
|
2873
|
+
# vs the mean beta of the other topics, on the full genes_clean basis.
|
|
2874
|
+
log2fc_final = _topic_log2fc(beta_final)
|
|
2875
|
+
topic_dict = {}
|
|
2876
|
+
for k in range(n_topics):
|
|
2877
|
+
top_idx = log2fc_final[k].argsort()[::-1][:15]
|
|
2878
|
+
topic_dict[f"Topic_{k}"] = [genes_s[i] for i in top_idx]
|
|
2879
|
+
qc_df = pd.DataFrame(topic_dict)
|
|
2880
|
+
|
|
2881
|
+
duration = time.perf_counter() - start_time
|
|
2882
|
+
print(f"Step 6 done in {duration:.2f}s")
|
|
2883
|
+
return Model(
|
|
2884
|
+
n_topics=n_topics,
|
|
2885
|
+
odg_names=list(odg_pack.odg_names),
|
|
2886
|
+
beta_consensus=beta_consensus,
|
|
2887
|
+
qc_df=qc_df,
|
|
2888
|
+
per_slice_betas=per_slice_betas_full,
|
|
2889
|
+
per_slice_stability=dict(per_slice_stability) if per_slice_stability else None,
|
|
2890
|
+
)
|
|
2891
|
+
|
|
2892
|
+
|
|
2893
|
+
# STEP 7: sampling engine (per-slice, with rescue)
|
|
2894
|
+
|
|
2895
|
+
@_profile_step
|
|
2896
|
+
def step7_sampling_engine(
|
|
2897
|
+
slices: Dict[str, SliceData],
|
|
2898
|
+
model: Model,
|
|
2899
|
+
config_path: Optional[str] = None,
|
|
2900
|
+
species: str = "hs",
|
|
2901
|
+
low_slice_quality=False,
|
|
2902
|
+
min_topic_percentage: Optional[float] = None,
|
|
2903
|
+
) -> Dict[str, dict]:
|
|
2904
|
+
"""Per-slice hierarchical Bayesian sampling. Returns dict of per-slice cell outputs.
|
|
2905
|
+
|
|
2906
|
+
Threshold-and-renormalize is always applied: any topic with theta < min_topic_percentage
|
|
2907
|
+
is zeroed and theta is renormalized. low_slice_quality=True per slice additionally
|
|
2908
|
+
enforces an inflate rescue: every surviving topic gets at least 1 cell.
|
|
2909
|
+
|
|
2910
|
+
The inner gene-and-cell allocation is vectorized via a sequential-binomial
|
|
2911
|
+
batched multinomial sampler, so the draws are statistically equivalent to
|
|
2912
|
+
the legacy per-gene-per-cell loop but the RNG draw order differs: outputs
|
|
2913
|
+
will not be bit-for-bit identical to the legacy implementation under the
|
|
2914
|
+
same seed. The marginals (per-spot UMI conservation, per-cell topic
|
|
2915
|
+
assignment, gamma cell weights) are preserved.
|
|
2916
|
+
"""
|
|
2917
|
+
start_time = time.perf_counter()
|
|
2918
|
+
profile = _load_config(config_path, species)
|
|
2919
|
+
cfg_min_pct = profile.get('min_topic_percentage', 0.05)
|
|
2920
|
+
if min_topic_percentage is None:
|
|
2921
|
+
min_topic_percentage = cfg_min_pct
|
|
2922
|
+
|
|
2923
|
+
names = list(slices.keys())
|
|
2924
|
+
low_q_d = _broadcast(low_slice_quality, names, default=False)
|
|
2925
|
+
|
|
2926
|
+
out: Dict[str, dict] = {}
|
|
2927
|
+
# Use the legacy RandomState (Mersenne Twister) so that multinomial / gamma
|
|
2928
|
+
# draws match the legacy codeconv Step 7 sequence sample-for-sample under the
|
|
2929
|
+
# same seed. PCG64 (np.random.default_rng) would draw a different sequence.
|
|
2930
|
+
rng = np.random.RandomState(_SEED)
|
|
2931
|
+
|
|
2932
|
+
for name in names:
|
|
2933
|
+
sd = slices[name]
|
|
2934
|
+
slice_low_q = low_q_d[name]
|
|
2935
|
+
engine_params = sd.engine_params
|
|
2936
|
+
gamma_shape = 1.0 / engine_params['phi']
|
|
2937
|
+
gamma_scale = engine_params['mu'] * engine_params['phi']
|
|
2938
|
+
|
|
2939
|
+
n_spots, n_genes = sd.counts_clean.shape
|
|
2940
|
+
n_topics = model.n_topics
|
|
2941
|
+
beta = sd.beta_final # (K, n_genes_s)
|
|
2942
|
+
theta = sd.theta # (n_spots, K)
|
|
2943
|
+
counts_csr = sd.counts_clean.tocsr()
|
|
2944
|
+
|
|
2945
|
+
# Triplet buffers built in chunks per spot; concatenated at the end.
|
|
2946
|
+
rows_chunks: List[np.ndarray] = []
|
|
2947
|
+
cols_chunks: List[np.ndarray] = []
|
|
2948
|
+
data_chunks: List[np.ndarray] = []
|
|
2949
|
+
cell_metadata = []
|
|
2950
|
+
global_cell_idx = 0
|
|
2951
|
+
n_rescued = 0
|
|
2952
|
+
|
|
2953
|
+
print(f"\nStep 7 [{name}]: sampling low_q={slice_low_q} min_pct={min_topic_percentage}")
|
|
2954
|
+
for s in tqdm(range(n_spots), desc=f"[{name}] cells", unit="spot"):
|
|
2955
|
+
n_total = int(sd.n_cells[s])
|
|
2956
|
+
if n_total == 0:
|
|
2957
|
+
continue
|
|
2958
|
+
|
|
2959
|
+
# Threshold-and-renormalize theta. Bypassed entirely when
|
|
2960
|
+
# min_topic_percentage <= 0 so theta_eff stays identical to theta[s].
|
|
2961
|
+
if min_topic_percentage > 0:
|
|
2962
|
+
theta_eff = theta[s].copy()
|
|
2963
|
+
theta_eff[theta_eff < min_topic_percentage] = 0.0
|
|
2964
|
+
tsum = theta_eff.sum()
|
|
2965
|
+
if tsum <= 0:
|
|
2966
|
+
theta_eff = np.zeros_like(theta_eff)
|
|
2967
|
+
theta_eff[int(np.argmax(theta[s]))] = 1.0
|
|
2968
|
+
tsum = 1.0
|
|
2969
|
+
theta_eff /= tsum
|
|
2970
|
+
else:
|
|
2971
|
+
theta_eff = theta[s]
|
|
2972
|
+
|
|
2973
|
+
topic_dist = _safe_multinomial(rng, n_total, theta_eff)
|
|
2974
|
+
|
|
2975
|
+
# Inflate rescue for low-quality slices.
|
|
2976
|
+
if slice_low_q:
|
|
2977
|
+
surviving = np.where(theta_eff > 0)[0]
|
|
2978
|
+
for k in surviving:
|
|
2979
|
+
if topic_dist[k] == 0:
|
|
2980
|
+
topic_dist[k] = 1
|
|
2981
|
+
n_rescued += 1
|
|
2982
|
+
|
|
2983
|
+
# Per-topic cell groups: draw gamma cell weights, register the
|
|
2984
|
+
# in-silico cell metadata, record the global cell indices.
|
|
2985
|
+
cells_weights: Dict[int, np.ndarray] = {}
|
|
2986
|
+
cells_indices: Dict[int, np.ndarray] = {}
|
|
2987
|
+
current_spot_cell_idx = 0
|
|
2988
|
+
for k in range(n_topics):
|
|
2989
|
+
n_k = int(topic_dist[k])
|
|
2990
|
+
if n_k <= 0:
|
|
2991
|
+
continue
|
|
2992
|
+
w = rng.gamma(gamma_shape, gamma_scale, size=n_k)
|
|
2993
|
+
if w.sum() == 0:
|
|
2994
|
+
w = np.ones(n_k)
|
|
2995
|
+
cells_weights[k] = w / w.sum()
|
|
2996
|
+
for i in range(n_k):
|
|
2997
|
+
cell_metadata.append({
|
|
2998
|
+
'spot_idx': s,
|
|
2999
|
+
'topic_idx': k,
|
|
3000
|
+
'cell_num': current_spot_cell_idx + i + 1,
|
|
3001
|
+
})
|
|
3002
|
+
cells_indices[k] = np.arange(global_cell_idx, global_cell_idx + n_k, dtype=np.int64)
|
|
3003
|
+
global_cell_idx += n_k
|
|
3004
|
+
current_spot_cell_idx += n_k
|
|
3005
|
+
|
|
3006
|
+
# Pull the spot's nonzero (gene, count) entries directly from CSR -
|
|
3007
|
+
# no full-row toarray(), no per-gene Python loop.
|
|
3008
|
+
start, end = counts_csr.indptr[s], counts_csr.indptr[s + 1]
|
|
3009
|
+
if start == end:
|
|
3010
|
+
continue
|
|
3011
|
+
gene_idx_arr = counts_csr.indices[start:end]
|
|
3012
|
+
gene_count_arr = counts_csr.data[start:end].astype(np.int64)
|
|
3013
|
+
n_present = gene_idx_arr.shape[0]
|
|
3014
|
+
|
|
3015
|
+
# p_topic for every (gene, topic) in one matmul. Rows that sum to
|
|
3016
|
+
# zero (no beta mass for any topic) fall back to uniform, matching
|
|
3017
|
+
# the legacy per-gene uniform fallback.
|
|
3018
|
+
beta_sub = beta[:, gene_idx_arr].T # (n_present, K)
|
|
3019
|
+
p_topic_mat = beta_sub * theta_eff[None, :]
|
|
3020
|
+
row_sums = p_topic_mat.sum(axis=1)
|
|
3021
|
+
zero_rows = row_sums <= 0
|
|
3022
|
+
if zero_rows.any():
|
|
3023
|
+
p_topic_mat[zero_rows] = 1.0 / n_topics
|
|
3024
|
+
row_sums[zero_rows] = 1.0
|
|
3025
|
+
p_topic_mat = p_topic_mat / row_sums[:, None]
|
|
3026
|
+
|
|
3027
|
+
# Batched UMI-per-topic for every present gene at once.
|
|
3028
|
+
umi_per_topic_mat = _batched_multinomial(rng, gene_count_arr, p_topic_mat)
|
|
3029
|
+
# umi_per_topic_mat: (n_present, K)
|
|
3030
|
+
|
|
3031
|
+
# Per topic, batch-allocate the topic UMIs across the topic's cells.
|
|
3032
|
+
for k in range(n_topics):
|
|
3033
|
+
if k not in cells_indices:
|
|
3034
|
+
continue
|
|
3035
|
+
n_k = cells_indices[k].shape[0]
|
|
3036
|
+
u_counts_k = umi_per_topic_mat[:, k]
|
|
3037
|
+
nonzero_g = np.flatnonzero(u_counts_k > 0)
|
|
3038
|
+
if nonzero_g.size == 0:
|
|
3039
|
+
continue
|
|
3040
|
+
# Broadcast cell weights across the genes-with-UMIs in this topic.
|
|
3041
|
+
w_tiled = np.broadcast_to(
|
|
3042
|
+
cells_weights[k][None, :], (nonzero_g.size, n_k)
|
|
3043
|
+
).copy()
|
|
3044
|
+
cell_alloc = _batched_multinomial(
|
|
3045
|
+
rng, u_counts_k[nonzero_g], w_tiled
|
|
3046
|
+
) # (n_genes_with_umi, n_k)
|
|
3047
|
+
|
|
3048
|
+
# Emit (row, col, data) triplets via np.where on the (gene, cell)
|
|
3049
|
+
# block, no Python-side cell loop.
|
|
3050
|
+
gi_idx, ci_idx = np.where(cell_alloc > 0)
|
|
3051
|
+
if gi_idx.size == 0:
|
|
3052
|
+
continue
|
|
3053
|
+
rows_chunks.append(cells_indices[k][ci_idx])
|
|
3054
|
+
cols_chunks.append(gene_idx_arr[nonzero_g[gi_idx]])
|
|
3055
|
+
data_chunks.append(cell_alloc[gi_idx, ci_idx].astype(np.int64))
|
|
3056
|
+
|
|
3057
|
+
if rows_chunks:
|
|
3058
|
+
rows = np.concatenate(rows_chunks)
|
|
3059
|
+
cols = np.concatenate(cols_chunks)
|
|
3060
|
+
data = np.concatenate(data_chunks)
|
|
3061
|
+
else:
|
|
3062
|
+
rows = np.zeros(0, dtype=np.int64)
|
|
3063
|
+
cols = np.zeros(0, dtype=np.int64)
|
|
3064
|
+
data = np.zeros(0, dtype=np.int64)
|
|
3065
|
+
|
|
3066
|
+
final_matrix = sp.csr_matrix((data, (rows, cols)), shape=(global_cell_idx, n_genes))
|
|
3067
|
+
print(f" [{name}] generated {global_cell_idx} cells rescued {n_rescued} topic slots")
|
|
3068
|
+
out[name] = {
|
|
3069
|
+
'matrix': final_matrix,
|
|
3070
|
+
'cell_metadata': cell_metadata,
|
|
3071
|
+
'n_cells': global_cell_idx,
|
|
3072
|
+
'n_rescued': n_rescued,
|
|
3073
|
+
}
|
|
3074
|
+
|
|
3075
|
+
duration = time.perf_counter() - start_time
|
|
3076
|
+
print(f"\nStep 7 done in {duration:.2f}s")
|
|
3077
|
+
return out
|
|
3078
|
+
|
|
3079
|
+
|
|
3080
|
+
# STEP 8: placement (per-slice)
|
|
3081
|
+
|
|
3082
|
+
# STEP 8: placement (per-slice)
|
|
3083
|
+
|
|
3084
|
+
@_profile_step
|
|
3085
|
+
def step8_geometry_and_placement(
|
|
3086
|
+
cells: Dict[str, dict],
|
|
3087
|
+
slices: Dict[str, SliceData],
|
|
3088
|
+
):
|
|
3089
|
+
"""Place each slice's in-silico cells around their spot centers using Fibonacci shells."""
|
|
3090
|
+
start_time = time.perf_counter()
|
|
3091
|
+
phi_golden = np.pi * (3. - np.sqrt(5.))
|
|
3092
|
+
|
|
3093
|
+
for name, payload in cells.items():
|
|
3094
|
+
sd = slices[name]
|
|
3095
|
+
diameter = sd.scale_factors['spot_diameter_fullres']
|
|
3096
|
+
cell_metadata = payload['cell_metadata']
|
|
3097
|
+
|
|
3098
|
+
spot_totals: Dict[int, int] = {}
|
|
3099
|
+
for meta in cell_metadata:
|
|
3100
|
+
spot_totals[meta['spot_idx']] = spot_totals.get(meta['spot_idx'], 0) + 1
|
|
3101
|
+
|
|
3102
|
+
final_barcodes: List[str] = []
|
|
3103
|
+
final_coords: List[List[float]] = []
|
|
3104
|
+
spot_counters: Dict[int, int] = {}
|
|
3105
|
+
|
|
3106
|
+
for meta in cell_metadata:
|
|
3107
|
+
s = meta['spot_idx']
|
|
3108
|
+
t = meta['topic_idx']
|
|
3109
|
+
idx = spot_counters.get(s, 0)
|
|
3110
|
+
spot_counters[s] = idx + 1
|
|
3111
|
+
n_total = spot_totals[s]
|
|
3112
|
+
|
|
3113
|
+
cy = sd.coords[s, 0]
|
|
3114
|
+
cx = sd.coords[s, 1]
|
|
3115
|
+
# Vogel sunflower confined to the spot footprint: the max radius is the
|
|
3116
|
+
# spot radius (diameter / 2), NOT the diameter, otherwise cells reach a
|
|
3117
|
+
# full diameter out and overlap neighbouring spots. (idx + 0.5) / n gives
|
|
3118
|
+
# an area-uniform fill of the disk.
|
|
3119
|
+
r = 0.5 * diameter * np.sqrt((idx + 0.5) / n_total) if n_total > 1 else 0.0
|
|
3120
|
+
ang = idx * phi_golden
|
|
3121
|
+
py = cy + r * np.sin(ang)
|
|
3122
|
+
px = cx + r * np.cos(ang)
|
|
3123
|
+
ay, ax = sd.coords[s, 2], sd.coords[s, 3]
|
|
3124
|
+
|
|
3125
|
+
final_coords.append([py, px, ay, ax])
|
|
3126
|
+
orig_bc = sd.barcodes[s]
|
|
3127
|
+
new_bc = f"{orig_bc}-Topic-{t}-Cell-{meta['cell_num']}"
|
|
3128
|
+
final_barcodes.append(new_bc)
|
|
3129
|
+
|
|
3130
|
+
payload['final_barcodes'] = final_barcodes
|
|
3131
|
+
payload['final_coords'] = np.array(final_coords)
|
|
3132
|
+
print(f"Step 8 [{name}]: placed {len(final_barcodes)} cells")
|
|
3133
|
+
|
|
3134
|
+
duration = time.perf_counter() - start_time
|
|
3135
|
+
print(f"Step 8 done in {duration:.2f}s")
|
|
3136
|
+
|
|
3137
|
+
|
|
3138
|
+
# STEP 9: export to 10X-compatible format (per-slice)
|
|
3139
|
+
|
|
3140
|
+
@_profile_step
|
|
3141
|
+
def step9_export_results(
|
|
3142
|
+
cells: Dict[str, dict],
|
|
3143
|
+
slices: Dict[str, SliceData],
|
|
3144
|
+
output_folder: str,
|
|
3145
|
+
spot_pitch_um: float = 100.0,
|
|
3146
|
+
interaction_range_um: float = 50.0,
|
|
3147
|
+
):
|
|
3148
|
+
"""Per-slice export to slice_<name>/deconvolved/ in 10X SpaceRanger HDF5 layout.
|
|
3149
|
+
|
|
3150
|
+
For a single slice run, the slice_<name> wrapper is still applied so behavior is
|
|
3151
|
+
consistent across single and multi-slice modes.
|
|
3152
|
+
|
|
3153
|
+
The per-slice ``run_summary.json`` carries a ``cellchat`` block with everything
|
|
3154
|
+
the downstream R session needs, ``ratio``, ``tol``, ``interaction_range``,
|
|
3155
|
+
``contact_range``, so the CellChat step reads one file instead of
|
|
3156
|
+
re-deriving unit conversions by hand::
|
|
3157
|
+
|
|
3158
|
+
cc <- jsonlite::fromJSON("slice_<name>/run_summary.json")$cellchat
|
|
3159
|
+
spatial.factors <- data.frame(ratio = cc$ratio, tol = cc$tol)
|
|
3160
|
+
|
|
3161
|
+
Parameters
|
|
3162
|
+
----------
|
|
3163
|
+
spot_pitch_um : float
|
|
3164
|
+
Centre-to-centre spot distance in microns, used to calibrate the
|
|
3165
|
+
pixel-to-micron ratio. 100.0 for all 10x Visium capture areas;
|
|
3166
|
+
pass the array pitch for other platforms.
|
|
3167
|
+
interaction_range_um : float
|
|
3168
|
+
The CellChat ``interaction.range`` to record and to run the sub-spot
|
|
3169
|
+
diagnostics against. Default 50.0, the shortest range that still covers
|
|
3170
|
+
the whole ~60 um Visium spot footprint, so every intra-spot pair is
|
|
3171
|
+
counted and the inference stays invariant to the Step 8 packing.
|
|
3172
|
+
See :func:`cellchat_spatial_factors`.
|
|
3173
|
+
"""
|
|
3174
|
+
start_time = time.perf_counter()
|
|
3175
|
+
os.makedirs(output_folder, exist_ok=True)
|
|
3176
|
+
|
|
3177
|
+
summary = {}
|
|
3178
|
+
for name, payload in cells.items():
|
|
3179
|
+
sd = slices[name]
|
|
3180
|
+
slice_dir = os.path.join(output_folder, f"slice_{name}")
|
|
3181
|
+
out_dir = os.path.join(slice_dir, "deconvolved")
|
|
3182
|
+
spa_dir = os.path.join(out_dir, "spatial")
|
|
3183
|
+
os.makedirs(out_dir, exist_ok=True)
|
|
3184
|
+
os.makedirs(spa_dir, exist_ok=True)
|
|
3185
|
+
|
|
3186
|
+
h5_path = os.path.join(out_dir, "filtered_feature_bc_matrix.h5")
|
|
3187
|
+
print(f"Step 9 [{name}]: writing {h5_path}")
|
|
3188
|
+
genes = sd.genes_clean
|
|
3189
|
+
barcodes = payload['final_barcodes']
|
|
3190
|
+
final_matrix = payload['matrix']
|
|
3191
|
+
|
|
3192
|
+
with h5py.File(h5_path, "w") as f:
|
|
3193
|
+
grp = f.create_group("matrix")
|
|
3194
|
+
grp.create_dataset("features/id", data=np.array(genes, dtype='S'))
|
|
3195
|
+
grp.create_dataset("features/name", data=np.array(genes, dtype='S'))
|
|
3196
|
+
grp.create_dataset("features/feature_type",
|
|
3197
|
+
data=np.array(["Gene Expression"] * len(genes), dtype='S'))
|
|
3198
|
+
grp.create_dataset("features/genome", data=np.array(["Genome"] * len(genes), dtype='S'))
|
|
3199
|
+
grp.create_dataset("features/_all_tag_keys", data=np.array([], dtype='S'))
|
|
3200
|
+
grp.create_dataset("barcodes", data=np.array(barcodes, dtype='S'))
|
|
3201
|
+
csc = final_matrix.T.tocsc()
|
|
3202
|
+
grp.create_dataset("data", data=csc.data)
|
|
3203
|
+
grp.create_dataset("indices", data=csc.indices)
|
|
3204
|
+
grp.create_dataset("indptr", data=csc.indptr)
|
|
3205
|
+
grp.create_dataset("shape", data=np.array(csc.shape, dtype='i4'))
|
|
3206
|
+
|
|
3207
|
+
# Copy spatial assets
|
|
3208
|
+
for fname in ['tissue_lowres_image.png', 'tissue_hires_image.png', 'scalefactors_json.json']:
|
|
3209
|
+
src = os.path.join(sd.spatial_path, fname)
|
|
3210
|
+
dst = os.path.join(spa_dir, fname)
|
|
3211
|
+
if os.path.exists(src):
|
|
3212
|
+
shutil.copy(src, dst)
|
|
3213
|
+
|
|
3214
|
+
coords = payload['final_coords']
|
|
3215
|
+
|
|
3216
|
+
# No registered histology (DBiT-seq samples vary): write a black
|
|
3217
|
+
# placeholder so Load10X_Spatial and every spatial plot still work.
|
|
3218
|
+
if not any(os.path.exists(os.path.join(spa_dir, f))
|
|
3219
|
+
for f in ('tissue_hires_image.png', 'tissue_lowres_image.png')):
|
|
3220
|
+
dummy_sf = make_dummy_tissue_image(
|
|
3221
|
+
coords, spa_dir,
|
|
3222
|
+
spot_diameter_fullres=sd.scale_factors.get('spot_diameter_fullres'),
|
|
3223
|
+
)
|
|
3224
|
+
merged = dict(sd.scale_factors or {})
|
|
3225
|
+
merged.update(dummy_sf)
|
|
3226
|
+
with open(os.path.join(spa_dir, 'scalefactors_json.json'), 'w') as fsf:
|
|
3227
|
+
json.dump(merged, fsf, indent=2)
|
|
3228
|
+
pos_df = pd.DataFrame(coords, columns=['pxl_row', 'pxl_col', 'array_row', 'array_col'])
|
|
3229
|
+
pos_df['barcode'] = barcodes
|
|
3230
|
+
pos_df['in_tissue'] = 1
|
|
3231
|
+
pos_df = pos_df[['barcode', 'in_tissue', 'array_row', 'array_col', 'pxl_row', 'pxl_col']]
|
|
3232
|
+
pos_df.to_csv(os.path.join(spa_dir, "tissue_positions.csv"), index=False)
|
|
3233
|
+
|
|
3234
|
+
# Per-slice summary. The CellChat block is recorded here so the downstream
|
|
3235
|
+
# R session can read spatial.factors straight out of run_summary.json
|
|
3236
|
+
# instead of guessing them (see cellchat_spatial_factors).
|
|
3237
|
+
slice_summary = {
|
|
3238
|
+
'slice_name': name,
|
|
3239
|
+
'n_spots': int(sd.counts_clean.shape[0]),
|
|
3240
|
+
'n_cells_generated': int(payload['n_cells']),
|
|
3241
|
+
'n_topics': int(sd.theta.shape[1]) if sd.theta is not None else None,
|
|
3242
|
+
'n_rescued_topic_slots': int(payload.get('n_rescued', 0)),
|
|
3243
|
+
'gene_count_export': int(len(genes)),
|
|
3244
|
+
}
|
|
3245
|
+
try:
|
|
3246
|
+
cc = cellchat_spatial_factors(
|
|
3247
|
+
{name: payload}, {name: sd},
|
|
3248
|
+
spot_pitch_um=spot_pitch_um,
|
|
3249
|
+
interaction_range_um=interaction_range_um,
|
|
3250
|
+
)[name]
|
|
3251
|
+
slice_summary['cellchat'] = {
|
|
3252
|
+
'ratio': round(cc['ratio'], 6),
|
|
3253
|
+
'tol': round(cc['tol'], 4),
|
|
3254
|
+
'interaction_range': round(float(interaction_range_um), 2),
|
|
3255
|
+
'contact_range': round(float(interaction_range_um), 2),
|
|
3256
|
+
'spot_pitch_um': float(spot_pitch_um),
|
|
3257
|
+
'spot_pitch_px': round(cc['spot_pitch_px'], 3),
|
|
3258
|
+
'spot_diameter_um': round(cc['spot_diameter_um'], 3),
|
|
3259
|
+
'min_intercell_distance_um': round(cc['min_distance_um'], 4),
|
|
3260
|
+
'median_nn_distance_um': round(cc['median_nn_distance_um'], 4),
|
|
3261
|
+
'mean_neighbours': round(cc.get('mean_neighbours', float('nan')), 3),
|
|
3262
|
+
'same_spot_fraction': round(cc.get('same_spot_fraction', float('nan')), 4),
|
|
3263
|
+
'intra_spot_pairs_captured': round(cc.get('intra_spot_pairs_captured', float('nan')), 4),
|
|
3264
|
+
}
|
|
3265
|
+
except Exception as exc: # never let a diagnostic break the export
|
|
3266
|
+
print(f" [{name}] could not derive CellChat spatial factors: {exc}")
|
|
3267
|
+
with open(os.path.join(slice_dir, "run_summary.json"), 'w') as fout:
|
|
3268
|
+
json.dump(slice_summary, fout, indent=2)
|
|
3269
|
+
summary[name] = slice_summary
|
|
3270
|
+
print(f" [{name}] cells={slice_summary['n_cells_generated']} genes={slice_summary['gene_count_export']}")
|
|
3271
|
+
|
|
3272
|
+
# Top-level summary across all slices
|
|
3273
|
+
with open(os.path.join(output_folder, "run_summary.json"), 'w') as fout:
|
|
3274
|
+
json.dump({'slices': summary, 'n_slices': len(summary)}, fout, indent=2)
|
|
3275
|
+
|
|
3276
|
+
duration = time.perf_counter() - start_time
|
|
3277
|
+
print(f"\nStep 9 done in {duration:.2f}s exports ready for Seurat::Load10X_Spatial()")
|