vima-spatial 0.1.0__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.
@@ -0,0 +1,40 @@
1
+ Metadata-Version: 2.2
2
+ Name: vima-spatial
3
+ Version: 0.1.0
4
+ Summary: variationa inference-based microniche analysis
5
+ Home-page: https://github.com/yakirr/vima
6
+ Author: Yakir Reshef
7
+ Author-email: yreshef@broadinstitute.org
8
+ Project-URL: Bug Tracker, https://github.com/yakirr/vima/issues
9
+ Project-URL: Tutorial, https://github.com/yakirr/vima/blob/main/README.md
10
+ Classifier: Programming Language :: Python :: 3
11
+ Classifier: License :: OSI Approved :: MIT License
12
+ Classifier: Operating System :: OS Independent
13
+ Requires-Python: >=3.12.3
14
+ Description-Content-Type: text/markdown
15
+ Requires-Dist: torch>=2.3.0
16
+ Requires-Dist: torchvision>=0.18.0
17
+ Requires-Dist: anndata>=0.10.7
18
+ Requires-Dist: matplotlib
19
+ Requires-Dist: multianndata
20
+ Requires-Dist: numpy>=1.26.4
21
+ Requires-Dist: scanpy
22
+ Requires-Dist: cv2
23
+ Requires-Dist: xarray
24
+ Requires-Dist: seaborn
25
+ Requires-Dist: pandas>=2.2.3
26
+ Requires-Dist: scipy
27
+ Requires-Dist: mpl_toolkits
28
+ Requires-Dist: cna>=0.1.7
29
+ Requires-Dist: tqdm
30
+ Requires-Dist: skimage
31
+ Requires-Dist: IPython
32
+
33
+ # vima
34
+ Variational inference-based microniche analysis is a method for conducting case-control analysis on multi-sample spatial molecular datasets. `vima` can be applied to any spatially resolved molecular technology, is well powered even at the modest sample sizes typical of research cohorts, and avoids traditional, parameter-intensive preprocessing steps such as cell segmentation or clustering of cells into discrete cell types. It works by treating each spatial sample as an image and using a variational autoencoder to extract numerical "fingerprints" from small tissue patches that capture their biological content. It uses these fingerprints to define a large number of "microniches'' – small, potentially overlapping groups of tissue patches with highly similar biology that span multiple samples. It then uses rigorous statistics to identify microniches whose abundance correlates with case-control status.
35
+
36
+ ## installation
37
+ To use `vima`, please clone this repository and add it to your `PYTHONPATH`. You will first need to install `pytorch'.
38
+
39
+ ## demo
40
+ Coming soon!
@@ -0,0 +1,8 @@
1
+ # vima
2
+ Variational inference-based microniche analysis is a method for conducting case-control analysis on multi-sample spatial molecular datasets. `vima` can be applied to any spatially resolved molecular technology, is well powered even at the modest sample sizes typical of research cohorts, and avoids traditional, parameter-intensive preprocessing steps such as cell segmentation or clustering of cells into discrete cell types. It works by treating each spatial sample as an image and using a variational autoencoder to extract numerical "fingerprints" from small tissue patches that capture their biological content. It uses these fingerprints to define a large number of "microniches'' – small, potentially overlapping groups of tissue patches with highly similar biology that span multiple samples. It then uses rigorous statistics to identify microniches whose abundance correlates with case-control status.
3
+
4
+ ## installation
5
+ To use `vima`, please clone this repository and add it to your `PYTHONPATH`. You will first need to install `pytorch'.
6
+
7
+ ## demo
8
+ Coming soon!
@@ -0,0 +1,6 @@
1
+ [build-system]
2
+ requires = [
3
+ "setuptools>=42",
4
+ "wheel"
5
+ ]
6
+ build-backend = "setuptools.build_meta"
@@ -0,0 +1,48 @@
1
+ [metadata]
2
+ name = vima-spatial
3
+ version = 0.1.0
4
+ author = Yakir Reshef
5
+ author_email = yreshef@broadinstitute.org
6
+ description = variationa inference-based microniche analysis
7
+ long_description = file: README.md
8
+ long_description_content_type = text/markdown
9
+ url = https://github.com/yakirr/vima
10
+ project_urls =
11
+ Bug Tracker = https://github.com/yakirr/vima/issues
12
+ Tutorial = https://github.com/yakirr/vima/blob/main/README.md
13
+ classifiers =
14
+ Programming Language :: Python :: 3
15
+ License :: OSI Approved :: MIT License
16
+ Operating System :: OS Independent
17
+
18
+ [options]
19
+ package_dir =
20
+ = src
21
+ packages = find:
22
+ python_requires = >=3.12.3
23
+ install_requires =
24
+ torch>=2.3.0
25
+ torchvision>=0.18.0
26
+ anndata>=0.10.7
27
+ matplotlib
28
+ multianndata
29
+ numpy>=1.26.4
30
+ scanpy
31
+ cv2
32
+ xarray
33
+ seaborn
34
+ pandas>=2.2.3
35
+ scipy
36
+ mpl_toolkits
37
+ cna>=0.1.7
38
+ tqdm
39
+ skimage
40
+ IPython
41
+
42
+ [options.packages.find]
43
+ where = src
44
+
45
+ [egg_info]
46
+ tag_build =
47
+ tag_date = 0
48
+
@@ -0,0 +1,2 @@
1
+ from . import data
2
+ from . import models
@@ -0,0 +1,77 @@
1
+ import numpy as np
2
+ import scanpy as sc
3
+ import anndata as ad
4
+ import pandas as pd
5
+ import multianndata as md
6
+ import torch
7
+ from torch.utils.data import DataLoader
8
+ import cna
9
+ from tqdm import tqdm
10
+ pb = lambda x: tqdm(x, ncols=100)
11
+
12
+ def anndata(patchmeta, Z, samplemeta, var_names=None, use_rep='X', n_comps=10, sampleid='sid'):
13
+ d = ad.AnnData(Z)
14
+ if var_names is not None:
15
+ d.var_names = var_names
16
+ d.obs = patchmeta
17
+
18
+ if use_rep == 'X_pca':
19
+ sc.tl.pca(d, n_comps=min(n_comps, Z.shape[1]-1))
20
+
21
+ print('running UMAP')
22
+ sc.pp.neighbors(d, use_rep=use_rep)
23
+ sc.tl.umap(d)
24
+
25
+ samplemeta.index = samplemeta.index.astype(str)
26
+ d.obs.sid = d.obs.sid.astype(str)
27
+ d = md.MultiAnnData(d)
28
+ d.samplem = samplemeta
29
+ d.sampleid = sampleid
30
+ print(f'built MultiAnnData object with {sampleid} as the unit of analysis')
31
+
32
+ return d
33
+
34
+ def apply(model, P, embedding=None, batch_size=1000):
35
+ if embedding is None:
36
+ embedding = model.embedding
37
+
38
+ P.pytorch_mode()
39
+ P.augmentation_off()
40
+ model.eval()
41
+ eval_loader = DataLoader(
42
+ dataset=P,
43
+ batch_size=batch_size,
44
+ shuffle=False)
45
+
46
+ Z = []
47
+ with torch.no_grad():
48
+ for batch in pb(eval_loader):
49
+ Z.append(embedding(batch).detach().cpu().numpy())
50
+
51
+ return np.concatenate(Z)
52
+
53
+ def latentrep(model, P, samplemeta):
54
+ return anndata(P.meta,
55
+ apply(model, P),
56
+ samplemeta[samplemeta.index.isin(P.meta.sid.unique())],
57
+ sampleid='sid')
58
+
59
+ def association(d, pheno, fdr=0.1, force_recompute=True, covs=None, Nnull=100000, seed=0, **kwargs):
60
+ if seed is not None: np.random.seed(seed)
61
+ cna.tl.nam(d, force_recompute=force_recompute)
62
+ d.samplem['case'] = pheno
63
+ if covs is not None:
64
+ d.samplem[covs.columns] = covs
65
+
66
+ res = cna.tl.association(d, d.samplem.case, donorids=d.samplem.donor.values, covs=covs, Nnull=Nnull, **kwargs)
67
+ print(f'P = {res.p}, used {res.k} MAM-PCs')
68
+ d.obs['mncoeff'] = res.ncorrs
69
+ if res.fdrs.fdr.min() <= fdr:
70
+ print(f'Found {res.fdrs[res.fdrs.fdr < 0.1].iloc[0].num_detected} microniches at FDR {int(fdr*100)}%')
71
+ d.obs['sig_mncoeff'] = res.ncorrs * (np.abs(res.ncorrs) > res.fdrs[res.fdrs.fdr < fdr].iloc[0].threshold)
72
+ else:
73
+ print(f'No microniches found at FDR {int(fdr*100)}%')
74
+ d.obs['sig_mncoeff'] = 0
75
+ d.samplem.loc[~np.isnan(d.samplem.case), 'yhat'] = res.yresid_hat
76
+
77
+ return res
@@ -0,0 +1,3 @@
1
+ from . import patchcollection
2
+ from . import samples
3
+ from . import ingest
@@ -0,0 +1,441 @@
1
+ import numpy as numpy
2
+ import pandas as pd
3
+ import numpy as np
4
+ import anndata as ad
5
+ import scanpy as sc
6
+ import xarray as xr
7
+ import cv2 as cv2
8
+ from skimage.filters import threshold_otsu
9
+ import seaborn as sns
10
+ import matplotlib.pyplot as plt
11
+ import matplotlib.colors as mcolors
12
+ import gc, os, subprocess
13
+ from tqdm import tqdm
14
+ pb = lambda x: tqdm(x, ncols=100)
15
+
16
+ compression = {'zlib': True, 'complevel': 2} # settings for writing xarrays
17
+
18
+ ###########################################
19
+ # utility functions
20
+ ###########################################
21
+ def xr_to_pixellist(s, mask):
22
+ return s.data[mask.data]
23
+
24
+ def set_pixels(s, mask, pl):
25
+ s.data[mask.data] = pl
26
+
27
+ def ar():
28
+ plt.gca().set_aspect('equal')
29
+
30
+ ###########################################
31
+ # for creating raw pixel files
32
+ ###########################################
33
+ def transcriptlist_to_pixellist(transcriptlist, x_colname='global_x', y_colname='global_y', gene_colname='gene', pixel_size=10):
34
+ # adds dummy rows such that there is at least one entry for every possible x- and y- value
35
+ # between the min and max values
36
+ def complete(pl, colname, genes, fill=0., verbose=True):
37
+ vals = np.sort(pl[colname].unique())
38
+ min_col = vals.min() // 1
39
+ max_col = vals.max() // 1
40
+ delta = int(min(vals[1:] - vals[:-1]))
41
+ full_range = list(np.arange(min_col, max_col + 1, delta))
42
+ locs_toadd = np.setdiff1d(full_range, vals)
43
+ if verbose: print(f'\tadding {colname}={locs_toadd}')
44
+ toadd = pl.iloc[:len(locs_toadd)].copy()
45
+ toadd[colname] = locs_toadd
46
+ toadd[genes] = fill
47
+ return pd.concat([pl, toadd], axis=0, ignore_index=True)
48
+
49
+ transcriptlist = transcriptlist[[x_colname, y_colname, gene_colname]].copy()
50
+ transcriptlist['pixel_x'] = (transcriptlist[x_colname] / pixel_size).astype(int) * pixel_size
51
+ transcriptlist['pixel_y'] = (transcriptlist[y_colname] / pixel_size).astype(int) * pixel_size
52
+
53
+ pixels = transcriptlist.groupby(['pixel_x', 'pixel_y'])[gene_colname].value_counts().unstack(fill_value=0)
54
+ pixels.reset_index(inplace=True)
55
+ pl = pixels.rename_axis(None, axis=1)
56
+ genes = pl.columns[2:]
57
+
58
+ return complete(complete(pl, 'pixel_x', genes), 'pixel_y', genes)
59
+
60
+ def pixellist_to_pixelmatrix(pl, markers):
61
+ # pivot in pandas
62
+ s = pd.pivot_table(pl, values=markers, index='pixel_y', columns='pixel_x').fillna(0)
63
+ s.columns.names = ['markers', 'pixel_x']
64
+
65
+ # convert to xarray
66
+ s = df_to_xarray32(s)
67
+ print('sample shape:', s.shape)
68
+
69
+ return s
70
+
71
+ def df_to_xarray32(df):
72
+ markers = df.columns.get_level_values('markers').unique()
73
+ return xr.DataArray(
74
+ df.values.reshape((len(df), len(markers), -1)).transpose(0,2,1),
75
+ coords={'x': df.columns.get_level_values('pixel_x').unique().values, 'y': df.index.values, 'marker': markers.values},
76
+ dims=['y', 'x', 'marker']
77
+ ).astype(np.float32)
78
+
79
+ def downsample(sample, factor, aggregate=np.mean):
80
+ pad_width = (
81
+ (int(factor - sample.shape[0] % factor), 0),
82
+ (int(factor - sample.shape[1] % factor), 0),
83
+ (0,0))
84
+ sample = np.pad(sample, pad_width, mode='constant', constant_values=0)
85
+ smaller = sample.reshape(sample.shape[0], sample.shape[1]//factor, factor, sample.shape[2])
86
+ smaller = aggregate(smaller, axis=2)
87
+ smaller = smaller.reshape(smaller.shape[0]//factor, factor, smaller.shape[1], smaller.shape[2])
88
+ smaller = aggregate(smaller, axis=1)
89
+ return smaller
90
+
91
+ def hiresarray_to_downsampledxarray(sample, name, factor, pixelsize, markers):
92
+ sample = downsample(sample, factor)
93
+ sample = xr.DataArray(
94
+ sample,
95
+ coords={'x': np.arange(sample.shape[1])*factor*pixelsize, 'y': np.arange(sample.shape[0])*factor*pixelsize, 'marker': markers},
96
+ dims=['y', 'x', 'marker']
97
+ ).astype(np.float32)
98
+ sample.name = name
99
+ return sample
100
+
101
+ ###########################################
102
+ # processing raw pixel files
103
+ ###########################################
104
+ def foreground_mask_st(s, min_ntranscripts=10):
105
+ totals = s.sum(dim='marker')
106
+ mask = totals > min_ntranscripts
107
+ return mask
108
+
109
+ def foreground_mask_ihc(s, real_markers, neg_ctrls, not_imaged_thresh, artifact_thresh, transform=lambda x:x, thresholding_method=threshold_otsu,
110
+ neg_ctrl_pseudocount=0, blur_width=5):
111
+ totals = (s.sel(marker=real_markers).sum(dim='marker') / (s.sel(marker=neg_ctrls).sum(dim='marker') + len(neg_ctrls) + neg_ctrl_pseudocount))
112
+ totals = transform(cv2.GaussianBlur(totals.data, (blur_width, blur_width),0))
113
+ valid_pixels = totals[(totals > not_imaged_thresh) & (totals < artifact_thresh)]
114
+ t = thresholding_method(valid_pixels)
115
+
116
+ return xr.DataArray(((totals > t) & (totals < artifact_thresh)).astype('bool'),
117
+ coords={'x': s.x, 'y': s.y},
118
+ dims=['y','x'], name=s.name)
119
+
120
+ def foreground_mask_codex(s, real_markers, neg_ctrls, blur_width=5):
121
+ # compute totals
122
+ totals = s.sel(marker=real_markers).sum(dim='marker')
123
+ totals = np.log1p(totals)
124
+ totals -= totals.min()
125
+ totals /= (totals.max()/255)
126
+ totals = totals.astype('uint16')
127
+
128
+ # determine foreground vs background
129
+ blurred = cv2.GaussianBlur(totals.data,(blur_width, blur_width),0)
130
+ _, mask = cv2.threshold(blurred,0,255,cv2.THRESH_BINARY+cv2.THRESH_OTSU)
131
+ return xr.DataArray(mask.astype('bool'),
132
+ coords={'x': totals.x, 'y': totals.y},
133
+ dims=['y','x'], name=s.name)
134
+
135
+ def write_masks(pixelsdir, outdir, get_foreground, sids, plot=True, vmax=30):
136
+ for sid in sids:
137
+ print('reading', sid)
138
+ s = xr.load_dataarray(f'{pixelsdir}/{sid}.nc').astype(np.float32)
139
+
140
+ # make mask and save
141
+ mask = get_foreground(s)
142
+ print(f'{mask.values.sum()} of {mask.shape[0]*mask.shape[1]} ({100*mask.values.sum()/(mask.shape[0]*mask.shape[1]):.0f}%) pixels are non-empty')
143
+ mask.to_netcdf(f'{outdir}/{sid}.nc', encoding={mask.name: compression}, engine="netcdf4")
144
+
145
+ if plot:
146
+ s.sum(dim='marker').plot(cmap='Reds', vmin=0, vmax=vmax); ar()
147
+ mask.plot(alpha=0.5, vmin=0, vmax=1, cmap='gray', add_colorbar=False)
148
+ plt.show()
149
+
150
+ subset = s.where(mask, other=0).sel(marker=s.marker[::10])
151
+ norm = mcolors.Normalize(vmin=subset.data.min(), vmax=0.95*subset.data.max())
152
+ subset.plot(col='marker', col_wrap=4, norm=norm)
153
+ plt.show()
154
+
155
+ gc.collect()
156
+
157
+ def get_sumstats_st(pixels):
158
+ ntranscripts = pixels.sum(axis=1, dtype=np.float64)
159
+ med_ntranscripts = np.median(ntranscripts)
160
+ pixels = np.log1p(med_ntranscripts * pixels / ntranscripts[:,None])
161
+ means = pixels.mean(axis=0, dtype=np.float64)
162
+ stds = pixels.std(axis=0, dtype=np.float64)
163
+ return {'means':means, 'stds':stds, 'med_ntranscripts':med_ntranscripts}
164
+
165
+ def normalize_st(mask, s, med_ntranscripts=None, means=None, stds=None):
166
+ s = s.where(mask, other=0)
167
+ pl = xr_to_pixellist(s, mask)
168
+ pl = np.log1p(med_ntranscripts * pl / pl.sum(axis=1)[:,None])
169
+ pl -= means
170
+ pl /= stds
171
+ set_pixels(s, mask, pl)
172
+ s.attrs['med_ntranscripts'] = med_ntranscripts
173
+ s.attrs['means'] = means
174
+ s.attrs['stds'] = stds
175
+ return s
176
+
177
+ def normalize_allsamples(pixelsdir, masksdir, outdir, sids, get_sumstats=get_sumstats_st, normalize=normalize_st):
178
+ print('reading all non-empty pixels')
179
+ pixels = np.concatenate([
180
+ xr_to_pixellist(
181
+ xr.open_dataarray(f'{pixelsdir}/{sid}.nc').astype(np.float32),
182
+ xr.open_dataarray(f'{masksdir}/{sid}.nc')
183
+ )
184
+ for sid in pb(sids)])
185
+ gc.collect()
186
+
187
+ print('computing sumstats')
188
+ sumstats = get_sumstats(pixels)
189
+ del pixels; gc.collect()
190
+
191
+ print('normalizing and writing')
192
+ for sid in pb(sids):
193
+ s = normalize(
194
+ xr.open_dataarray(f'{masksdir}/{sid}.nc'),
195
+ xr.open_dataarray(f'{pixelsdir}/{sid}.nc').astype(np.float32),
196
+ **sumstats)
197
+ s.to_netcdf(f'{outdir}/{sid}.nc', encoding={s.name: compression}, engine="netcdf4")
198
+
199
+ ###########################################
200
+ # dimensionality reduction and integration
201
+ ###########################################
202
+ def metapixels_allsamples(normedpixelsdir, masksdir, sids, plot=True, ncols=8):
203
+ def cdf(v, ax):
204
+ sorted_data = np.sort(v)
205
+ cdf = np.arange(1, len(sorted_data) + 1) / len(sorted_data)
206
+ ax.plot(sorted_data, cdf)
207
+
208
+ all_metapixels = {}
209
+ all_npixels = {}
210
+
211
+ if plot:
212
+ nrows = int(np.ceil(len(sids)/ncols))
213
+ fig, axs = plt.subplots(nrows, ncols, figsize=(2*ncols,1.5*nrows))
214
+ axs = axs.reshape((nrows, -1))
215
+
216
+ for i, sid in enumerate(sids):
217
+ print('.', end='')
218
+ all_metapixels[sid], all_npixels[sid] = metapixels(
219
+ xr.open_dataarray(f'{normedpixelsdir}/{sid}.nc').astype(np.float32),
220
+ xr.open_dataarray(f'{masksdir}/{sid}.nc'))
221
+
222
+ # visualize distribution of num non-empty pixels per metapixel in this sample
223
+ if plot:
224
+ ax = axs[i // ncols, i % ncols]
225
+ cdf(all_npixels[sid], ax)
226
+ ax.set_title(sid)
227
+ gc.collect()
228
+
229
+ if plot:
230
+ plt.tight_layout()
231
+ plt.show()
232
+
233
+ return all_metapixels, all_npixels
234
+
235
+ def metapixels(s, mask, npixels_thresh=0):
236
+ markers = s.marker.values
237
+
238
+ # make metapixels and compute how many non-empty pixels and transcripts are in each metapixel
239
+ kernel = np.ones((5,5),np.float32)
240
+ mp = cv2.filter2D(s.data, -1, kernel)
241
+ npixels = cv2.filter2D(mask.data.astype('float32'), -1, kernel)
242
+
243
+ # filter out metapixels with few non-empty pixels
244
+ metapixels_mask = npixels > npixels_thresh
245
+
246
+ # divide each metapixel by the # of non-empty pixels that contributed to it and return
247
+ return pd.DataFrame(data=mp[metapixels_mask] / npixels[metapixels_mask][:,None], columns=markers), npixels[metapixels_mask]
248
+
249
+ # mps should be an array of dataframes containing metapixels
250
+ def pca_metapixels(mps, k, plot=True):
251
+ print('merging and standardizing metapixels')
252
+ allmp = pd.concat(mps)
253
+ allmp -= allmp.values.mean(axis=0, dtype=np.float64)
254
+ allmp /= allmp.values.std(axis=0, dtype=np.float64)
255
+ allmp = ad.AnnData(X=allmp)
256
+ C = np.corrcoef(allmp.X[::max(1,(len(allmp)//50000))].T)
257
+ print(allmp.shape)
258
+
259
+ print('performing PCA')
260
+ sc.tl.pca(allmp, n_comps=k)
261
+ loadings = pd.DataFrame(data=allmp.varm['PCs'], columns=[f'PC{i}' for i in range(1,k+1)], index=allmp.var_names)
262
+
263
+ if plot:
264
+ plt.imshow(C, cmap='seismic', vmin=-1, vmax=1)
265
+ plt.show()
266
+ plt.figure(figsize=(30,2))
267
+ plt.imshow(loadings.T, cmap='seismic', vmin=-0.5, vmax=0.5)
268
+ plt.xticks(range(len(loadings)), loadings.index, rotation=90)
269
+ plt.show()
270
+
271
+ return loadings, C, allmp
272
+
273
+ def pca_pixels(normedixelsdir, masksdir, pcloadings, sids, plot=True, npixels_to_plot=50000, colorby=['sid']):
274
+ print('reading in pixels')
275
+ pls = np.concatenate([
276
+ xr_to_pixellist(
277
+ xr.open_dataarray(f'{normedixelsdir}/{sid}.nc').astype(np.float32),
278
+ xr.open_dataarray(f'{masksdir}/{sid}.nc'))
279
+ for sid in pb(sids)])
280
+ sid_labels = np.concatenate([
281
+ np.array([sid] * xr.open_dataarray(f'{masksdir}/{sid}.nc').sum().item())
282
+ for sid in sids
283
+ ])
284
+ print('applying dimensionality reduction')
285
+ allpixels_pca = pd.DataFrame(
286
+ pls.dot(pcloadings),
287
+ columns=[f'PC{i}' for i in range(1,pcloadings.shape[1]+1)]
288
+ )
289
+ allpixels_pca['sid'] = sid_labels
290
+ del pls; gc.collect()
291
+
292
+ if plot:
293
+ print('visualizing')
294
+ visualize_pixels(allpixels_pca, npixels_to_plot, colorby)
295
+
296
+ return allpixels_pca
297
+
298
+ def harmonize(allpixels_pca, outdir, integrate=['sid']):
299
+ path_to_data = os.path.abspath(f'{outdir}/_allpixels_pca.feather')
300
+ path_to_script = os.path.dirname(__file__) + '/harmonize.R'
301
+ command = ['Rscript', path_to_script, path_to_data] + integrate
302
+ print('Please run the following command in your R environment:')
303
+ print(' '.join(command))
304
+ print()
305
+ print('When this finishes, run vi.post_harmony')
306
+
307
+ def visualize_pixels(pixels, ntoplot, colorby, include_pca_plot=False):
308
+ pcs = [c for c in pixels.columns if c.startswith('PC')]
309
+ metavars = [c for c in pixels.columns if c not in pcs]
310
+ np.random.seed(0)
311
+ ix = np.random.choice(len(pixels), replace=False, size=ntoplot)
312
+ toplot = pixels.iloc[ix]
313
+ toplot_ad = ad.AnnData(
314
+ X=toplot[pcs],
315
+ obs=toplot[metavars])
316
+ sc.pp.neighbors(toplot_ad, use_rep='X')
317
+ sc.tl.umap(toplot_ad)
318
+
319
+ for metavar in colorby:
320
+ if include_pca_plot:
321
+ sns.scatterplot(x='PC1', y='PC2', hue=metavar, data=toplot, palette='Set1', s=1, legend=False)
322
+ plt.title(metavar)
323
+ plt.show()
324
+ sc.pl.umap(toplot_ad, color=metavar)
325
+
326
+ return toplot_ad
327
+
328
+ def write_harmonized(masksdir, outdir, harmpixels, sids):
329
+ pcs = [c for c in harmpixels.columns if c.startswith('PC')]
330
+ hpcs = ['h'+c for c in pcs]
331
+ for sid in pb(sids):
332
+ mask = xr.open_dataarray(f'{masksdir}/{sid}.nc')
333
+ pl = harmpixels[harmpixels.sid == sid]
334
+ s_ = np.zeros((*mask.shape, len(hpcs)))
335
+ s_[mask.data] = pl[pcs].values
336
+ s = xr.DataArray(s_,
337
+ dims=['y', 'x', 'marker'],
338
+ coords={'x': mask.x, 'y': mask.y, 'marker': hpcs})
339
+ s.name = sid
340
+ s.to_netcdf(f'{outdir}/{sid}.nc', encoding={s.name: compression}, engine="netcdf4")
341
+ gc.collect()
342
+
343
+ ###########################################
344
+ # user-facing interface
345
+ ###########################################
346
+ import glob
347
+ def preprocess(outdir, repname, get_foreground, get_sumstats, normalize,
348
+ sid_to_covs=None, nmetamarkers=10, plot=False):
349
+ # prepare directory structure
350
+ countsdir = f'{outdir}/counts'
351
+ normeddir = f'{outdir}/normalized'
352
+ masksdir = f'{outdir}/masks'
353
+ processeddir = f'{outdir}/{repname}'
354
+ os.makedirs(normeddir, exist_ok=True)
355
+ os.makedirs(masksdir, exist_ok=True)
356
+ os.makedirs(processeddir, exist_ok=True)
357
+
358
+ # prepare
359
+ sids = [f.split('/')[-1].split('.nc')[0]
360
+ for f in glob.glob(f'{countsdir}/*.nc')]
361
+ if sid_to_covs is not None:
362
+ harmony_cov_names = list(sid_to_covs.columns)
363
+ else:
364
+ harmony_cov_names = []
365
+
366
+ # create and write masks
367
+ write_masks(countsdir, masksdir, get_foreground, sids, plot=plot)
368
+
369
+ # create normalized pixels
370
+ normalize_allsamples(countsdir, masksdir, normeddir, sids,
371
+ get_sumstats=get_sumstats,
372
+ normalize=normalize)
373
+
374
+ # create metapixels for more accurate PCA
375
+ metapixels, npixels = metapixels_allsamples(normeddir, masksdir, sids, plot=plot)
376
+
377
+ # PCA the metapixels
378
+ loadings, C, allmp = pca_metapixels(metapixels.values(), nmetamarkers)
379
+ loadings.to_feather(f'{processeddir}/_pcloadings.feather')
380
+ del metapixels, allmp; gc.collect()
381
+
382
+ # apply the PC loadings to plain pixels
383
+ allpixels_pca = pca_pixels(normeddir, masksdir, loadings, sids)
384
+
385
+ for cov_name in harmony_cov_names:
386
+ allpixels_pca[cov_name] = allpixels_pca['sid'].map(sid_to_covs[cov_name])
387
+ allpixels_pca.to_feather(f'{processeddir}/_allpixels_pca.feather')
388
+
389
+ # prompt user to run harmony
390
+ harmonize(allpixels_pca, processeddir, integrate=['sid'] + harmony_cov_names)
391
+
392
+ def post_harmony(outdir, repname, sid_to_covs=None, plot=True):
393
+ masksdir = f'{outdir}/masks'
394
+ processeddir = f'{outdir}/{repname}'
395
+ sids = [f.split('/')[-1].split('.nc')[0]
396
+ for f in glob.glob(f'{masksdir}/*.nc')]
397
+ if sid_to_covs is not None:
398
+ harmony_cov_names = list(sid_to_covs.columns)
399
+ else:
400
+ harmony_cov_names = []
401
+
402
+ # read in harmonized pixels, visualize, and write
403
+ harmpixels = pd.read_feather(f'{processeddir}/_allpixels_pca_harmony.feather')
404
+ if plot:
405
+ viz = visualize_pixels(harmpixels, 50000, ['sid'] + harmony_cov_names)
406
+ write_harmonized(masksdir, processeddir, harmpixels, sids)
407
+
408
+ def sanity_checks(outdir, repname, sid_to_covs=None):
409
+ processeddir = f'{outdir}/{repname}'
410
+ sids = [f.split('/')[-1].split('.nc')[0]
411
+ for f in glob.glob(f'{processeddir}/*.nc')]
412
+ if sid_to_covs is not None:
413
+ harmony_cov_names = list(sid_to_covs.columns)
414
+ else:
415
+ harmony_cov_names = []
416
+
417
+ print('all PCs of one sample')
418
+ s = xr.open_dataarray(f'{processeddir}/{sids[0]}.nc').astype(np.float32)
419
+ s.plot(col='marker', col_wrap=5, vmin=-10, vmax=10, cmap='seismic')
420
+
421
+ print('histogram of each pc')
422
+ harmpixels = pd.read_feather(f'{processeddir}/_allpixels_pca_harmony.feather')
423
+ nmms = harmpixels.values.shape[1] - len(harmony_cov_names) - 1 # the -1 accounts for sid
424
+ plt.figure(figsize=(3*4, 2*int(np.ceil(nmms/4))))
425
+ for i in range(nmms):
426
+ print(i, end='')
427
+ plt.subplot(int(np.ceil(nmms/4)), 4, i+1)
428
+ plt.hist(harmpixels.values[:,i], bins=1000)
429
+ plt.tight_layout()
430
+ plt.show()
431
+
432
+ print('PC1 of several samples')
433
+ fig, axs = plt.subplots(len(sids[::5])//5 + 1, 5, figsize=(16, 4*(len(sids[::5])//5 + 1)))
434
+ for sid, ax in zip(sids[::3], axs.flatten()):
435
+ s = xr.open_dataarray(f'{processeddir}/{sid}.nc').astype(np.float32)
436
+ vmax = np.percentile(np.abs(s.sel(marker='hPC1').data), 99)
437
+ s.sel(marker='hPC1').plot(ax=ax, cmap='seismic', vmin=-vmax, vmax=vmax, add_colorbar=False)
438
+ ax.set_title(sid)
439
+ gc.collect()
440
+ plt.tight_layout()
441
+ plt.show()