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()")