vima-spatial 0.2.0__tar.gz → 0.2.2__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (38) hide show
  1. {vima_spatial-0.2.0/src/vima_spatial.egg-info → vima_spatial-0.2.2}/PKG-INFO +2 -2
  2. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/setup.cfg +2 -2
  3. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/cc.py +30 -10
  4. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/fingerprints.py +11 -1
  5. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/ingest/dimreduce.py +22 -3
  6. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/ingest/ingest.py +2 -2
  7. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/ingest/st.py +7 -7
  8. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/patchfeatures.py +78 -55
  9. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/vis/__init__.py +1 -1
  10. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/vis/patches.py +27 -10
  11. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/vis/spatial.py +8 -6
  12. {vima_spatial-0.2.0 → vima_spatial-0.2.2/src/vima_spatial.egg-info}/PKG-INFO +2 -2
  13. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima_spatial.egg-info/requires.txt +1 -1
  14. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/MANIFEST.in +0 -0
  15. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/README.md +0 -0
  16. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/pyproject.toml +0 -0
  17. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/__init__.py +0 -0
  18. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/data/__init__.py +0 -0
  19. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/data/patchcollection.py +0 -0
  20. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/data/samples.py +0 -0
  21. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/ingest/__init__.py +0 -0
  22. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/ingest/nonst.py +0 -0
  23. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/ingest/util.py +0 -0
  24. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/models/__init__.py +0 -0
  25. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/models/resnet_vae.py +0 -0
  26. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/models/resnetlight_decoder.py +0 -0
  27. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/models/resnetlight_encoder.py +0 -0
  28. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/models/simple_vae.py +0 -0
  29. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/models/vae.py +0 -0
  30. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/train/__init__.py +0 -0
  31. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/train/logging.py +0 -0
  32. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/train/training.py +0 -0
  33. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/vis/features.py +0 -0
  34. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/vis/patchexamples.py +0 -0
  35. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima/vis/umaps.py +0 -0
  36. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima_spatial.egg-info/SOURCES.txt +0 -0
  37. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima_spatial.egg-info/dependency_links.txt +0 -0
  38. {vima_spatial-0.2.0 → vima_spatial-0.2.2}/src/vima_spatial.egg-info/top_level.txt +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: vima-spatial
3
- Version: 0.2.0
3
+ Version: 0.2.2
4
4
  Summary: variational inference-based microniche analysis
5
5
  Home-page: https://github.com/yakirr/vima
6
6
  Author: Yakir Reshef
@@ -25,7 +25,7 @@ Requires-Dist: netcdf4
25
25
  Requires-Dist: seaborn
26
26
  Requires-Dist: pandas>=2.2.3
27
27
  Requires-Dist: scipy
28
- Requires-Dist: cna>=0.2.3
28
+ Requires-Dist: cna>=0.2.4
29
29
  Requires-Dist: tqdm
30
30
  Requires-Dist: pyarrow
31
31
  Requires-Dist: scikit-image
@@ -1,6 +1,6 @@
1
1
  [metadata]
2
2
  name = vima-spatial
3
- version = 0.2.0
3
+ version = 0.2.2
4
4
  author = Yakir Reshef
5
5
  author_email = yreshef@broadinstitute.org
6
6
  description = variational inference-based microniche analysis
@@ -35,7 +35,7 @@ install_requires =
35
35
  seaborn
36
36
  pandas>=2.2.3
37
37
  scipy
38
- cna>=0.2.3
38
+ cna>=0.2.4
39
39
  tqdm
40
40
  pyarrow
41
41
  scikit-image
@@ -28,32 +28,52 @@ def anndata(patchmeta, Z, var_names=None, use_rep='X', n_comps=10, **kwargs):
28
28
 
29
29
  return d
30
30
 
31
- def apply(models, P, batch_size=1000):
31
+ def apply(models, P, batch_size=1000, with_mse=False):
32
32
  P.pytorch_mode()
33
33
  P.augmentation_off()
34
34
  for model in models:
35
35
  model.eval()
36
-
36
+
37
37
  eval_loader = DataLoader(
38
38
  dataset=P,
39
39
  batch_size=batch_size,
40
40
  shuffle=False)
41
41
 
42
42
  Zs = {modelid: [] for modelid in range(len(models))}
43
+ if with_mse:
44
+ MSEs = {modelid: [] for modelid in range(len(models))}
43
45
  with torch.no_grad():
44
46
  for batch in pb(eval_loader):
45
47
  for modelid, model in enumerate(models):
46
- Zs[modelid].append(model.embedding(batch).detach().cpu().numpy())
47
-
48
- return np.array([
49
- np.concatenate(Z) for Z in Zs.values()])
50
-
51
- def latentreps(models, P, use_rep='X', n_comps=100, **kwargs):
48
+ if with_mse:
49
+ x_recon, mean, _ = model(batch, sample_from_latent=False)
50
+ Zs[modelid].append(mean.reshape(len(batch[0]), -1).detach().cpu().numpy())
51
+ MSEs[modelid].append(((x_recon - batch[0]) ** 2).mean(dim=(2, 3)).detach().cpu().numpy())
52
+ else:
53
+ Zs[modelid].append(model.embedding(batch).detach().cpu().numpy())
54
+
55
+ Zs_out = np.array([np.concatenate(Z) for Z in Zs.values()])
56
+ if with_mse:
57
+ MSEs_out = np.array([np.concatenate(M) for M in MSEs.values()])
58
+ return Zs_out, MSEs_out
59
+ return Zs_out
60
+
61
+ def latentreps(models, P, use_rep='X', n_comps=100, with_mse=True, **kwargs):
52
62
  print('applying models')
53
- Zs = apply(models, P)
63
+ result = apply(models, P, with_mse=with_mse)
64
+ Zs, MSEs = result if with_mse else (result, None)
65
+
54
66
  print('computing nearest-neighbor graphs')
55
67
  ds = [anndata(P.meta, Z, use_rep=use_rep, n_comps=n_comps, **kwargs) for Z in pb(Zs)]
56
- return Fingerprints.from_list(ds)
68
+ fp = Fingerprints.from_list(ds)
69
+
70
+ if MSEs is not None:
71
+ mean_mse = MSEs.mean(axis=0)
72
+ marker_names = next(iter(P.samples.values())).coords['marker'].values
73
+ per_channel = pd.DataFrame(mean_mse, index=fp.obs.index, columns=marker_names)
74
+ fp.obsm['per_channel_mse'] = per_channel
75
+ fp.obs['mse'] = per_channel.mean(axis=1)
76
+ return fp
57
77
 
58
78
  def _association(MAMresid, M, Nmodels, y, batches, donorids, ks=None, Nnull=1000, show_progress=False):
59
79
  # prep data
@@ -120,7 +120,7 @@ class Fingerprints:
120
120
  self._adata.obsp[f'connectivities_{i}'] = d.obsp['connectivities']
121
121
  self._adata.obsp[f'distances_{i}'] = d.obsp['distances']
122
122
  self._adata.uns[f'neighbors_{i}'] = d.uns['neighbors']
123
-
123
+
124
124
  def sample_pcs(self, sid_name='sid'):
125
125
  D = self.avg_graph(make_umap=False)
126
126
  NAM, _ = cna.tl.nam(D, sid_name)
@@ -129,6 +129,16 @@ class Fingerprints:
129
129
  U, _, _ = np.linalg.svd(NAM, full_matrices=False)
130
130
  return pd.DataFrame(U, index=NAM.index,
131
131
  columns=[f'PC{i+1}' for i in range(U.shape[1])])
132
+
133
+ def mn_pcs(self, sid_name='sid'):
134
+ D = self.avg_graph(make_umap=False)
135
+ NAM, _ = cna.tl.nam(D, sid_name)
136
+ NAM -= NAM.mean(axis=0)
137
+ NAM /= NAM.std(axis=0)
138
+ _, _, VT = np.linalg.svd(NAM, full_matrices=False)
139
+ print(VT.shape)
140
+ return pd.DataFrame(VT.T, index=NAM.columns,
141
+ columns=[f'PC{i+1}' for i in range(VT.shape[0])])
132
142
 
133
143
  def to_anndata(self):
134
144
  X = np.hstack([self._adata.obsm[f'X_{i}'] for i in range(self.nmodels)])
@@ -25,16 +25,31 @@ def metapixels_allsamples(normedpixelsdir, masksdir, sids, total_n_metapixels, p
25
25
  # figure out how many metapixels to store per sample
26
26
  nsamples = len(sids)
27
27
  nmp_per_sample = total_n_metapixels // nsamples
28
-
29
28
 
30
29
  if plot:
31
30
  fig = plt.figure(figsize=(7,5))
32
31
 
33
32
  print('Creating metapixels prior to PCA')
34
33
  print(f'\t(will randomly downsample to {nmp_per_sample} metapixels per sample if needed.)')
34
+ ref_markers = None
35
+ ref_sid = None
35
36
  for i, sid in enumerate(pb(sids)):
36
37
  da = xr.open_dataarray(f'{normedpixelsdir}/{sid}.nc')
37
38
  mask_da = xr.open_dataarray(f'{masksdir}/{sid}.nc')
39
+
40
+ # ensure same markers in same order in all files
41
+ markers = list(da.marker.values)
42
+ if ref_markers is None:
43
+ ref_markers, ref_sid = markers, sid
44
+ elif markers != ref_markers:
45
+ missing = set(ref_markers) - set(markers)
46
+ extra = set(markers) - set(ref_markers)
47
+ order_mismatch = not missing and not extra
48
+ detail = (f'order differs' if order_mismatch else
49
+ f'{len(missing)} missing, {len(extra)} extra vs {ref_sid}')
50
+ print(f'\033[93mWARNING: {sid} has different markers ({len(markers)}) '
51
+ f'than {ref_sid} ({len(ref_markers)}): {detail}\033[0m')
52
+
38
53
  all_metapixels[sid], all_npixels[sid] = metapixels(da.astype(np.float32), mask_da)
39
54
  da.close(); mask_da.close()
40
55
  del da, mask_da
@@ -118,10 +133,14 @@ def pca_pixels(normedpixelsdir, masksdir, pcloadings, sids):
118
133
  for sid in pb(sids):
119
134
  da = xr.open_dataarray(f'{normedpixelsdir}/{sid}.nc')
120
135
  mask_da = xr.open_dataarray(f'{masksdir}/{sid}.nc')
121
- pl = util.xr_to_pixellist(da.astype(np.float32), mask_da)
122
- # close file handles before the dot product so HDF5 cache is released
136
+ # load raw arrays and close before dtype conversion so we never hold
137
+ # two full (H × W × n_genes) copies simultaneously
138
+ data = da.values
139
+ mask = mask_da.values
123
140
  da.close(); mask_da.close()
124
141
  del da, mask_da; gc.collect()
142
+ pl = data.astype(np.float32, copy=False)[mask]
143
+ del data, mask; gc.collect()
125
144
 
126
145
  pl_pca = pl.dot(pcloadings)
127
146
  pcs.append(pl_pca)
@@ -100,7 +100,7 @@ def pca_pixels(outdir, repname, nmetamarkers=10, plot=True, npixels_to_plot=5000
100
100
 
101
101
  if plot:
102
102
  visualize_pixels(pca, npixels_to_plot, 'metamarkers', cov_names)
103
- return pca
103
+ return pca, loadings
104
104
 
105
105
  def harmonize(allpixels_pca, sid_to_covs=None, npixels_to_plot=50000, plot=True):
106
106
  import harmonypy as hm
@@ -129,12 +129,12 @@ def write_harmonized(outdir, repname, harmpixels):
129
129
  pl = harmpixels[harmpixels.sid == sid]
130
130
  s_ = np.zeros((*mask.shape, len(hpcs)))
131
131
  s_[mask.data] = pl[pcs].values
132
- mask.close(); del mask
133
132
  s = xr.DataArray(s_,
134
133
  dims=['y', 'x', 'marker'],
135
134
  coords={'x': mask.x, 'y': mask.y, 'marker': hpcs})
136
135
  s.name = sid
137
136
  s.to_netcdf(f'{processeddir}/{sid}.nc', encoding={s.name: util.compression}, engine="netcdf4")
137
+ mask.close(); del mask
138
138
  gc.collect()
139
139
 
140
140
  def sanity_checks(outdir, repname, npcs=1, nskip=3):
@@ -171,7 +171,7 @@ def transcriptlist_to_normedpixelmatrix(sid, data, x_col, y_col, gene_col, pixel
171
171
 
172
172
  def rasterize_and_normalize_generic(load, filepaths, x_col, y_col, gene_col, n_top_genes_per_sample, pixel_size, outdir,
173
173
  min_ntranscripts_per_pixel, min_ngenes_per_pixel,
174
- genes_to_add=[], plot_mean_var=True, plot_spatial_hvgs=False):
174
+ genes_to_add=[], basic_plots=True, plot_spatial_hvgs=False):
175
175
  if len(filepaths) == 0:
176
176
  print('No files found. Check your filepaths and try again.')
177
177
  return
@@ -181,7 +181,7 @@ def rasterize_and_normalize_generic(load, filepaths, x_col, y_col, gene_col, n_t
181
181
  print('Finding HVGs and dataset-wide mean and variance per gene...')
182
182
  hvgs, means, stds = get_sumstats(load, filepaths, normfactor, x_col, y_col, gene_col,
183
183
  n_top_genes_per_sample=n_top_genes_per_sample,
184
- genes_to_add=genes_to_add, pixel_size=pixel_size, plot_mean_var=plot_mean_var,
184
+ genes_to_add=genes_to_add, pixel_size=pixel_size, plot_mean_var=basic_plots,
185
185
  plot_spatial_hvgs=plot_spatial_hvgs,
186
186
  min_ntranscripts=min_ntranscripts_per_pixel)
187
187
  print('Final number of genes used =', len(hvgs))
@@ -195,7 +195,7 @@ def rasterize_and_normalize_generic(load, filepaths, x_col, y_col, gene_col, n_t
195
195
  sid, data = load(filepath)
196
196
  print(f'Processing sample {i+1}/{len(filepaths)}: {sid}')
197
197
  mask, pm = transcriptlist_to_normedpixelmatrix(sid, data, x_col, y_col, gene_col, pixel_size,
198
- normfactor, means, stds, genes=hvgs, plots=plot_mean_var,
198
+ normfactor, means, stds, genes=hvgs, plots=basic_plots,
199
199
  min_ntranscripts_per_pixel=min_ntranscripts_per_pixel,
200
200
  min_ngenes_per_pixel=min_ngenes_per_pixel)
201
201
  del data; gc.collect()
@@ -203,25 +203,25 @@ def rasterize_and_normalize_generic(load, filepaths, x_col, y_col, gene_col, n_t
203
203
  util.write_xarray(pm, f'{normdir}/{pm.name}.nc')
204
204
 
205
205
  def prepare_xenium5k(load, filepaths, x_col, y_col, gene_col, n_top_genes_per_sample, outdir,
206
- pixel_size=10, genes_to_add=[], plot_mean_var=True, plot_spatial_hvgs=False,
206
+ pixel_size=10, genes_to_add=[], basic_plots=True, plot_spatial_hvgs=False,
207
207
  min_ntranscripts_per_pixel=11, min_ngenes_per_pixel=5):
208
208
  rasterize_and_normalize_generic(load, filepaths, x_col, y_col, gene_col,
209
209
  n_top_genes_per_sample,
210
210
  pixel_size=pixel_size,
211
211
  outdir=outdir,
212
212
  genes_to_add=genes_to_add,
213
- plot_mean_var=plot_mean_var,
213
+ basic_plots=basic_plots,
214
214
  plot_spatial_hvgs=plot_spatial_hvgs,
215
215
  min_ntranscripts_per_pixel=min_ntranscripts_per_pixel,
216
216
  min_ngenes_per_pixel=min_ngenes_per_pixel)
217
217
 
218
218
  def prepare_merfish(load, filepaths, x_col, y_col, gene_col, outdir,
219
- pixel_size=10, plot_mean_var=True,
219
+ pixel_size=10, basic_plots=True,
220
220
  min_ntranscripts_per_pixel=11, min_ngenes_per_pixel=1):
221
221
  rasterize_and_normalize_generic(load, filepaths, x_col, y_col, gene_col,
222
222
  None,
223
223
  pixel_size=pixel_size,
224
224
  outdir=outdir,
225
- plot_mean_var=plot_mean_var,
225
+ basic_plots=basic_plots,
226
226
  min_ntranscripts_per_pixel=min_ntranscripts_per_pixel,
227
227
  min_ngenes_per_pixel=min_ngenes_per_pixel)
@@ -4,7 +4,7 @@ import pandas as pd
4
4
  import xarray as xr
5
5
  from scipy.stats import rankdata
6
6
  from tqdm import tqdm
7
- pb = lambda x: tqdm(x, ncols=100)
7
+ pb = lambda x, desc: tqdm(x, ncols=100, desc=desc)
8
8
 
9
9
  def cell_type_counts(
10
10
  cells,
@@ -18,6 +18,7 @@ def cell_type_counts(
18
18
  patch_x_microns_col='x_microns',
19
19
  patch_y_microns_col='y_microns',
20
20
  patch_size_in_pixels_col='patchsize',
21
+ include_totalcells=False,
21
22
  pixel_size_microns=10
22
23
  ):
23
24
  """Return per-patch cell type counts.
@@ -46,7 +47,7 @@ def cell_type_counts(
46
47
  cell_types = sorted(cells[celltype_col].unique())
47
48
  counts = pd.DataFrame(0, index=patch_meta.index, columns=cell_types, dtype=int)
48
49
 
49
- for sid, sid_cells in pb(cells.groupby(sid_col)):
50
+ for sid, sid_cells in pb(cells.groupby(sid_col), 'cell_type_counts'):
50
51
  sid_patches = patch_meta[patch_meta[patch_sid_col] == sid]
51
52
  if len(sid_patches) == 0 or len(sid_cells) == 0:
52
53
  continue
@@ -70,6 +71,7 @@ def cell_type_counts(
70
71
  if normalized:
71
72
  totals = counts.sum(axis=1)
72
73
  counts = counts.div(totals, axis=0).fillna(0)
74
+ if include_totalcells:
73
75
  counts['totalcells'] = totals
74
76
 
75
77
  return counts
@@ -107,7 +109,7 @@ def expression_profiles(
107
109
 
108
110
  result = None
109
111
 
110
- for sid in pb(sids):
112
+ for sid in pb(sids, 'expression_profiles'):
111
113
  sid_patches = patch_meta[patch_meta[patch_sid_col] == sid]
112
114
  sample = xr.open_dataarray(os.path.join(directory, f'{sid}.nc')).load()
113
115
  marker_names = sample.coords['marker'].values.tolist()
@@ -137,45 +139,94 @@ def expression_profiles(
137
139
  return result
138
140
 
139
141
 
142
+ def _permutation_pvals(X_ranked, group_a, group_b, donors, n_perms, rng):
143
+ sum_a = X_ranked[group_a].sum(axis=0)
144
+ sum_b = X_ranked[group_b].sum(axis=0)
145
+ count_a = float(group_a.sum())
146
+ count_b = float(group_b.sum())
147
+ obs_diff = sum_a / count_a - sum_b / count_b
148
+
149
+ unique_donors = np.unique(donors)
150
+ n_donors = len(unique_donors)
151
+ da_sum = np.zeros((n_donors, X_ranked.shape[1]))
152
+ db_sum = np.zeros((n_donors, X_ranked.shape[1]))
153
+ da_count = np.zeros(n_donors)
154
+ db_count = np.zeros(n_donors)
155
+ for i, d in enumerate(unique_donors):
156
+ in_d = donors == d
157
+ da_sum[i] = X_ranked[in_d & group_a].sum(axis=0)
158
+ db_sum[i] = X_ranked[in_d & group_b].sum(axis=0)
159
+ da_count[i] = (in_d & group_a).sum()
160
+ db_count[i] = (in_d & group_b).sum()
161
+
162
+ delta_sum = db_sum - da_sum
163
+ delta_count = db_count - da_count
164
+ flip = (rng.random((n_perms, n_donors)) < 0.5).astype(float)
165
+
166
+ sum_a_null = sum_a + flip @ delta_sum
167
+ count_a_null = count_a + flip @ delta_count
168
+ sum_b_null = (sum_a + sum_b) - sum_a_null
169
+ count_b_null = (count_a + count_b) - count_a_null
170
+ count_a_null = np.maximum(count_a_null, 1.0)
171
+ count_b_null = np.maximum(count_b_null, 1.0)
172
+ null_diff = sum_a_null / count_a_null[:, None] - sum_b_null / count_b_null[:, None]
173
+
174
+ return ((np.abs(null_diff) >= np.abs(obs_diff)).sum(axis=0) + 1) / (n_perms + 1)
175
+
176
+
177
+ def _ttest_pvals(X_ranked, group_a, group_b, donors):
178
+ from scipy import stats
179
+ global_mean_a = X_ranked[group_a].mean(axis=0)
180
+ global_mean_b = X_ranked[group_b].mean(axis=0)
181
+ diffs = []
182
+ for d in np.unique(donors):
183
+ in_d = donors == d
184
+ a_rows = X_ranked[in_d & group_a]
185
+ b_rows = X_ranked[in_d & group_b]
186
+ mean_a = a_rows.mean(axis=0) if len(a_rows) > 0 else global_mean_a
187
+ mean_b = b_rows.mean(axis=0) if len(b_rows) > 0 else global_mean_b
188
+ diffs.append(mean_a - mean_b)
189
+ if len(diffs) < 2:
190
+ raise ValueError(f'T-test requires at least 2 units; got {len(diffs)}')
191
+ _, pvals = stats.ttest_1samp(np.array(diffs), 0, axis=0)
192
+ return pvals
193
+
194
+
140
195
  def test_features(
141
196
  features,
142
197
  group_a,
143
198
  group_b=None,
144
199
  *,
145
- perm_key,
146
- method='mean_of_ranks',
200
+ unit_of_analysis,
147
201
  n_perms=100000,
148
202
  seed=None,
149
203
  corr_method='benjamini-hochberg',
204
+ Ttest=False,
150
205
  ):
151
- """Compare feature distributions between two patch groups via donor-level permutation.
206
+ """Compare feature distributions between two patch groups.
152
207
 
153
- Group labels are flipped at the donor level so the null distribution respects
154
- within-donor correlations.
208
+ Uses mean rank difference (Wilcoxon/AUC equivalent) as the test statistic.
209
+ By default, significance is assessed via donor-level permutation (group labels
210
+ are flipped at the unit_of_analysis level). Pass Ttest=True to instead run a
211
+ paired T-test on per-unit mean rank differences.
155
212
 
156
213
  Args:
157
214
  features: DataFrame (n_patches × n_features), e.g. from cell_type_counts or
158
215
  expression_profiles.
159
216
  group_a: boolean array, length n_patches — first group (e.g. associated patches).
160
217
  group_b: boolean array or None — second group; defaults to ~group_a.
161
- perm_key: array-like of donor/sample IDs aligned with features rows. Permutations
162
- flip labels at the level of unique values in perm_key.
163
- method: test statistic to use —
164
- 'mean_of_ranks' mean rank difference (Wilcoxon/AUC equivalent, default),
165
- 'mean_of_medians' mean of per-donor median differences; only donors
166
- with patches in both groups contribute.
167
- n_perms: number of permutations (default 100000).
168
- seed: random seed for reproducibility.
218
+ unit_of_analysis: array-like of donor/sample IDs aligned with features rows.
219
+ Permutations (or T-test pairing) operate at the level of unique values.
220
+ n_perms: number of permutations (default 100000); ignored when Ttest=True.
221
+ seed: random seed for reproducibility; ignored when Ttest=True.
169
222
  corr_method: multiple-testing correction — 'benjamini-hochberg' (default) or 'bonferroni'.
223
+ Ttest: if True, use a one-sample two-sided T-test on per-unit mean rank differences
224
+ instead of permutation testing (default False).
170
225
 
171
226
  Returns:
172
227
  DataFrame indexed by feature name, columns: median_a, median_b, diff, pvals,
173
228
  pvals_adj. diff = median_a - median_b. Sorted by pvals ascending.
174
229
  """
175
- if method not in ('mean_of_medians', 'mean_of_ranks'):
176
- raise ValueError(
177
- f'method must be "mean_of_medians", or "mean_of_ranks"; got {method!r}'
178
- )
179
230
  if corr_method not in ('benjamini-hochberg', 'bonferroni'):
180
231
  raise ValueError(
181
232
  f'corr_method must be "benjamini-hochberg" or "bonferroni"; got {corr_method!r}'
@@ -196,7 +247,7 @@ def test_features(
196
247
 
197
248
  group_a = _align(group_a, 'group_a').astype(bool)
198
249
  group_b = ~group_a if group_b is None else _align(group_b, 'group_b').astype(bool)
199
- donors = _align(perm_key, 'perm_key')
250
+ donors = _align(unit_of_analysis, 'unit_of_analysis')
200
251
  X = features.values.astype(float)
201
252
 
202
253
  print(f'Comparing {group_a.sum()} patches (Group A) to {group_b.sum()} patches (Group B).')
@@ -206,41 +257,12 @@ def test_features(
206
257
  median_a_raw = np.median(X[group_a], axis=0)
207
258
  median_b_raw = np.median(X[group_b], axis=0)
208
259
 
209
- if method == 'mean_of_ranks':
210
- X = rankdata(X, axis=0, nan_policy='raise')
211
-
212
- sum_a = X[group_a].sum(axis=0)
213
- sum_b = X[group_b].sum(axis=0)
214
- count_a = float(group_a.sum())
215
- count_b = float(group_b.sum())
216
- obs_diff = sum_a / count_a - sum_b / count_b
217
-
218
- unique_donors = np.unique(donors)
219
- n_donors = len(unique_donors)
220
- da_sum = np.zeros((n_donors, X.shape[1]))
221
- db_sum = np.zeros((n_donors, X.shape[1]))
222
- da_count = np.zeros(n_donors)
223
- db_count = np.zeros(n_donors)
224
- for i, d in enumerate(unique_donors):
225
- in_d = donors == d
226
- da_sum[i] = X[in_d & group_a].sum(axis=0)
227
- db_sum[i] = X[in_d & group_b].sum(axis=0)
228
- da_count[i] = (in_d & group_a).sum()
229
- db_count[i] = (in_d & group_b).sum()
260
+ X = rankdata(X, axis=0, nan_policy='raise')
230
261
 
231
- delta_sum = db_sum - da_sum
232
- delta_count = db_count - da_count
233
- flip = (rng.random((n_perms, n_donors)) < 0.5).astype(float)
234
-
235
- sum_a_null = sum_a + flip @ delta_sum
236
- count_a_null = count_a + flip @ delta_count
237
- sum_b_null = (sum_a + sum_b) - sum_a_null
238
- count_b_null = (count_a + count_b) - count_a_null
239
- count_a_null = np.maximum(count_a_null, 1.0)
240
- count_b_null = np.maximum(count_b_null, 1.0)
241
- null_diff = sum_a_null / count_a_null[:, None] - sum_b_null / count_b_null[:, None]
242
-
243
- pvals = ((np.abs(null_diff) >= np.abs(obs_diff)).sum(axis=0) + 1) / (n_perms + 1)
262
+ if Ttest:
263
+ pvals = _ttest_pvals(X, group_a, group_b, donors)
264
+ else:
265
+ pvals = _permutation_pvals(X, group_a, group_b, donors, n_perms, rng)
244
266
 
245
267
  if corr_method == 'benjamini-hochberg':
246
268
  from statsmodels.stats.multitest import multipletests
@@ -260,6 +282,7 @@ def test_features(
260
282
  'mean_b': mean_b,
261
283
  'log2fc': np.log2((median_a_raw + 1e-9) / (median_b_raw + 1e-9)),
262
284
  'log2fc_means': np.log2((mean_a + 1e-9) / (mean_b + 1e-9)),
285
+ 'diff_median': median_a_raw - median_b_raw,
263
286
  }, index=features.columns)
264
287
  else:
265
288
  result = pd.DataFrame({
@@ -5,4 +5,4 @@ from .patchexamples import (scaler, apply_colormap, plot_with_reconstruction,
5
5
  plot_patches_separatechannels, plot_patches_overlaychannels,
6
6
  plot_patches_overlaychannels_linsum,
7
7
  plot_patches_overlaychannels_sorted, plot_patches_fourcolors)
8
- from .spatial import spatialplot, annotate_spatialplot
8
+ from .spatial import spatialplot, annotate_spatialplot, plot_npatches_per_sample, plot_samples_with_patches, plot_sample_with_patches
@@ -1,3 +1,4 @@
1
+ import os
1
2
  import warnings
2
3
  import matplotlib.pyplot as plt
3
4
  import matplotlib.patches as mpatches
@@ -34,7 +35,8 @@ def _plot_separate(patches, markers, vmin, vmax, cmap='seismic', show=True):
34
35
  return fig
35
36
 
36
37
 
37
- def _plot_composite(patches, markers, colors, vmin, vmax, features=None, nx=5, ny=5, show=True):
38
+ def _plot_composite(patches, markers, colors, vmin, vmax, features=None, nx=5, ny=5, show=True,
39
+ subfig=None):
38
40
  N, ps, K = patches.shape[0], patches.shape[1], len(markers)
39
41
 
40
42
  rgb = np.zeros((N, ps, ps, 3))
@@ -58,7 +60,8 @@ def _plot_composite(patches, markers, colors, vmin, vmax, features=None, nx=5, n
58
60
  else:
59
61
  cell_to_patch = {i: i for i in range(N)}
60
62
 
61
- fig, axs = plt.subplots(ny, nx, figsize=(nx, ny))
63
+ fig = subfig if subfig is not None else plt.figure(figsize=(nx, ny))
64
+ axs = fig.subplots(ny, nx)
62
65
  for ax in axs.flatten():
63
66
  ax.axis('off')
64
67
  for cell_i, patch_i in cell_to_patch.items():
@@ -68,7 +71,8 @@ def _plot_composite(patches, markers, colors, vmin, vmax, features=None, nx=5, n
68
71
  legend_handles = [mpatches.Patch(facecolor=colors[k], label=markers[k]) for k in range(K)]
69
72
  fig.legend(handles=legend_handles, loc='lower center', ncol=K, frameon=False,
70
73
  fontsize=8, bbox_to_anchor=(0.5, 0), bbox_transform=fig.transFigure)
71
- plt.tight_layout(rect=[0, 0.08, 1, 1])
74
+ if subfig is None:
75
+ plt.tight_layout(rect=[0, 0.08, 1, 1])
72
76
  if show:
73
77
  plt.show()
74
78
  return fig
@@ -102,6 +106,10 @@ class MarkersInSpace:
102
106
  self._arrays = {} # {sid: np.ndarray(H, W, K)}
103
107
  self.vmin = {} # {marker: float}
104
108
  self.vmax = {} # {marker: float}
109
+ if samples is not None:
110
+ self._all_sids = set(samples.keys())
111
+ else:
112
+ self._all_sids = {f[:-3] for f in os.listdir(directory) if f.endswith('.nc')}
105
113
  if markers:
106
114
  self.add_markers(markers)
107
115
 
@@ -144,6 +152,7 @@ class MarkersInSpace:
144
152
  for sid, arr in pb(self._arrays.items(), f'Adding {len(new_markers)} markers'):
145
153
  new_data = self._read_sid(sid, new_markers)
146
154
  self._arrays[sid] = np.concatenate([arr, new_data], axis=-1)
155
+ self._ensure_sids_loaded(self._all_sids) # load remaining sids for dataset-wide stats
147
156
  self._update_stats()
148
157
 
149
158
  def _ensure_sids_loaded(self, sids):
@@ -203,7 +212,7 @@ class MarkersInSpace:
203
212
  return _plot_separate(patches, markers, vmin, vmax, cmap, show=show)
204
213
 
205
214
  def show_composite(self, patchmeta, markers=None, features=None, colors=None,
206
- n=25, nx=5, ny=5, seed=None, vmin=None, vmax=None, show=True):
215
+ n=25, nx=5, ny=5, seed=None, vmin=None, vmax=None, show=True, subfig=None):
207
216
  """Show patches as additive RGB composites in an nx × ny grid.
208
217
 
209
218
  Args:
@@ -244,7 +253,8 @@ class MarkersInSpace:
244
253
  marker_indices = [self._marker_to_idx[m] for m in markers]
245
254
  patches = self._extract_patches(patchmeta, marker_indices)
246
255
  vmin, vmax = self._resolve_scale(markers, vmin, vmax)
247
- return _plot_composite(patches, markers, colors, vmin, vmax, features, nx, ny, show=show)
256
+ return _plot_composite(patches, markers, colors, vmin, vmax, features, nx, ny, show=show,
257
+ subfig=subfig)
248
258
 
249
259
 
250
260
  # ── standalone convenience functions (backed by a global MarkersInSpace) ─────
@@ -273,7 +283,8 @@ def show_patches_separate(patchmeta, markers, directory, samples=None,
273
283
 
274
284
  def show_patches_composite(patchmeta, markers, directory, samples=None,
275
285
  features=None, colors=None,
276
- n=25, nx=5, ny=5, seed=None, vmin=None, vmax=None, show=True):
286
+ n=25, nx=5, ny=5, seed=None, vmin=None, vmax=None, show=True,
287
+ subfig=None):
277
288
  """Convenience wrapper around MarkersInSpace.show_composite using a global cache.
278
289
 
279
290
  On the first call (or when directory changes) a new MarkersInSpace instance is
@@ -282,12 +293,13 @@ def show_patches_composite(patchmeta, markers, directory, samples=None,
282
293
  """
283
294
  return _get_default_mis(directory, samples).show_composite(
284
295
  patchmeta, markers, features=features, colors=colors,
285
- n=n, nx=nx, ny=ny, seed=seed, vmin=vmin, vmax=vmax, show=show)
296
+ n=n, nx=nx, ny=ny, seed=seed, vmin=vmin, vmax=vmax, show=show, subfig=subfig)
286
297
 
287
298
 
288
299
  def show_patches_cells(patchmeta, cells, x_col, y_col, celltype_col,
289
300
  pixelsize_microns, nx=5, ny=5, sid_col='sid',
290
- colors=None, seed=None, s=8, show=True):
301
+ colors=None, seed=None, s=8, show=True, subfig=None,
302
+ include_only=None):
291
303
  """Show an nx×ny grid of randomly chosen patches with cells overlaid as colored dots.
292
304
 
293
305
  Args:
@@ -312,13 +324,17 @@ def show_patches_cells(patchmeta, cells, x_col, y_col, celltype_col,
312
324
 
313
325
  extent = int(patchmeta['patchsize'].iloc[0]) * pixelsize_microns
314
326
 
327
+ if include_only is not None:
328
+ cells = cells[cells[celltype_col].isin(include_only)]
329
+
315
330
  # Build color map for cell types
316
331
  all_types = sorted(cells[celltype_col].dropna().unique())
317
332
  if colors is None:
318
333
  palette = plt.cm.tab20(np.linspace(0, 1, max(len(all_types), 1)))
319
334
  colors = {ct: palette[i] for i, ct in enumerate(all_types)}
320
335
 
321
- fig, axs = plt.subplots(ny, nx, figsize=(nx * 2, ny * 2), squeeze=False)
336
+ fig = subfig if subfig is not None else plt.figure(figsize=(nx * 2, ny * 2))
337
+ axs = fig.subplots(ny, nx, squeeze=False)
322
338
  for ax in axs.flatten():
323
339
  ax.set_visible(False)
324
340
 
@@ -352,7 +368,8 @@ def show_patches_cells(patchmeta, cells, x_col, y_col, celltype_col,
352
368
  fig.legend(handles=handles, loc='lower center', ncol=ncol_legend,
353
369
  frameon=False, fontsize=7,
354
370
  bbox_to_anchor=(0.5, 0), bbox_transform=fig.transFigure)
355
- plt.tight_layout(rect=[0, bottom_margin, 1, 1])
371
+ if subfig is None:
372
+ plt.tight_layout(rect=[0, bottom_margin, 1, 1])
356
373
 
357
374
  if show:
358
375
  plt.show()
@@ -129,14 +129,17 @@ def spatialplot(patchmeta, values, sids=None, cmap='viridis', vmin=None, vmax=No
129
129
  for ax in axs[len(sids):]:
130
130
  ax.axis('off')
131
131
 
132
+ fig._vima_sid_to_ax = sid_to_ax
132
133
  fig.tight_layout()
133
134
  if show:
134
135
  plt.show()
135
- else:
136
- return sid_to_ax
136
+ return fig
137
137
 
138
138
 
139
- def annotate_spatialplot(sid_to_ax, patchmeta, highlight, color, thickness=3, show=True):
139
+ def annotate_spatialplot(patchmeta, highlight, color, thickness=3, show=True, fig=None):
140
+ if fig is None:
141
+ fig = plt.gcf()
142
+ sid_to_ax = fig._vima_sid_to_ax
140
143
  for sid, ax in sid_to_ax.items():
141
144
  mypatches = patchmeta[patchmeta.sid == sid]
142
145
  if len(mypatches) == 0:
@@ -161,9 +164,8 @@ def annotate_spatialplot(sid_to_ax, patchmeta, highlight, color, thickness=3, sh
161
164
  cnt = cnt.squeeze()
162
165
  if cnt.ndim == 1:
163
166
  continue
164
- ax.plot(cnt[:, 0], cnt[:, 1], color=color, linewidth=thickness)
167
+ ax.plot(np.append(cnt[:, 0], cnt[0, 0]), np.append(cnt[:, 1], cnt[0, 1]), color=color, linewidth=thickness)
165
168
 
166
169
  if show:
167
170
  plt.show()
168
- else:
169
- return sid_to_ax
171
+ return fig
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: vima-spatial
3
- Version: 0.2.0
3
+ Version: 0.2.2
4
4
  Summary: variational inference-based microniche analysis
5
5
  Home-page: https://github.com/yakirr/vima
6
6
  Author: Yakir Reshef
@@ -25,7 +25,7 @@ Requires-Dist: netcdf4
25
25
  Requires-Dist: seaborn
26
26
  Requires-Dist: pandas>=2.2.3
27
27
  Requires-Dist: scipy
28
- Requires-Dist: cna>=0.2.3
28
+ Requires-Dist: cna>=0.2.4
29
29
  Requires-Dist: tqdm
30
30
  Requires-Dist: pyarrow
31
31
  Requires-Dist: scikit-image
@@ -11,7 +11,7 @@ netcdf4
11
11
  seaborn
12
12
  pandas>=2.2.3
13
13
  scipy
14
- cna>=0.2.3
14
+ cna>=0.2.4
15
15
  tqdm
16
16
  pyarrow
17
17
  scikit-image
File without changes
File without changes