vima-spatial 0.1.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.
vima/__init__.py ADDED
@@ -0,0 +1,2 @@
1
+ from . import data
2
+ from . import models
vima/association.py ADDED
@@ -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
vima/data/__init__.py ADDED
@@ -0,0 +1,3 @@
1
+ from . import patchcollection
2
+ from . import samples
3
+ from . import ingest
vima/data/ingest.py ADDED
@@ -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()
@@ -0,0 +1,125 @@
1
+ from torch.utils.data import Dataset
2
+ import torchvision.transforms.functional as TF
3
+ from torchvision import transforms
4
+ import numpy as np
5
+ import pandas as pd
6
+ import random
7
+ import torch
8
+ from . import samples as vds
9
+ from tqdm import tqdm
10
+ pb = lambda x: tqdm(x, ncols=100)
11
+
12
+ class ToTorch:
13
+ def __call__(self, x):
14
+ return torch.tensor(x).permute(*range(x.ndim - 3), x.ndim-1, x.ndim-3, x.ndim-2)
15
+
16
+ class RandomDiscreteRotation:
17
+ def __call__(self, x):
18
+ ntimes = np.random.choice([0,1,2,3])
19
+ for i in range(ntimes):
20
+ x = torch.rot90(x, dims=[-2,-1])
21
+ return x
22
+
23
+ class PatchCollection(Dataset):
24
+ @staticmethod
25
+ def choose_patches(samples, patchsize, patchstride, max_frac_empty):
26
+ patchmeta = []
27
+
28
+ for s in pb(samples.values()):
29
+ mask = vds.get_mask(s)
30
+ starts = np.array([
31
+ [i, j]
32
+ for i in range(0, mask.sizes['x']-patchsize, patchstride)
33
+ for j in range(0, mask.sizes['y']-patchsize, patchstride)
34
+ if mask.data[j:j+patchsize, i:i+patchsize].mean() > (1-max_frac_empty)
35
+ ]).astype('int')
36
+
37
+ patchmeta.append(pd.DataFrame([
38
+ (s.sid, s.donor, i, j, mask.x[i], mask.y[j])
39
+ for i, j in starts
40
+ ],
41
+ columns=['sid','donor','x','y', 'x_microns', 'y_microns'],
42
+ ))
43
+ patchmeta = pd.concat(patchmeta, axis=0).reset_index(drop=True)
44
+ patchmeta.x = patchmeta.x.astype('int')
45
+ patchmeta.y = patchmeta.y.astype('int')
46
+ patchmeta.x_microns = patchmeta.x_microns.astype('float32')
47
+ patchmeta.y_microns = patchmeta.y_microns.astype('float32')
48
+ patchmeta['patchsize'] = patchsize
49
+ return patchmeta
50
+
51
+ def __init__(self, samples, patchsize=40, patchstride=10, max_frac_empty=0.8,
52
+ sid_nums=None, standardize=True, percentile_thresh=99):
53
+ self.samples = samples
54
+ self.meta = PatchCollection.choose_patches(samples, patchsize, patchstride, max_frac_empty)
55
+ self.nmarkers = next(iter(samples.values())).sizes['marker']
56
+
57
+ self.pytorch_mode()
58
+ self.__preprocess__(standardize, percentile_thresh, sid_nums=sid_nums)
59
+ self.augmentation_off()
60
+
61
+ @property
62
+ def sid_nums(self):
63
+ return {sid:sid_num for sid, sid_num in self.meta[['sid','sid_num']].drop_duplicates().values}
64
+
65
+ @property
66
+ def nsamples(self):
67
+ return len(self.meta.sid.unique())
68
+
69
+ def augmentation_on(self):
70
+ if self.dim_order != 'pytorch':
71
+ print('Data augmentation only available in pytorch mode. Will leave augmentation off')
72
+ return
73
+ print('data augmentation is on')
74
+ self.transform = transforms.Compose([
75
+ ToTorch(),
76
+ RandomDiscreteRotation(),
77
+ transforms.RandomHorizontalFlip(),
78
+ ])
79
+ def augmentation_off(self):
80
+ print('data augmentation is off')
81
+ self.transform = transforms.Compose([
82
+ ToTorch(),
83
+ ])
84
+
85
+ def pytorch_mode(self):
86
+ self.dim_order = 'pytorch'
87
+ print('in pytorch mode')
88
+ def numpy_mode(self):
89
+ self.dim_order = 'numpy'
90
+ self.augmentation_off()
91
+ print('in numpy mode')
92
+
93
+ def __preprocess__(self, standardize, percentile_thresh, sid_nums=None):
94
+ self.patches = np.array([
95
+ self.samples[s].data[y:y+ps,x:x+ps,:]
96
+ for s, x, y, ps in self.meta[['sid','x','y','patchsize']].values
97
+ ])
98
+ if sid_nums is None:
99
+ self.meta['sid_num'] = pd.factorize(self.meta.sid)[0]
100
+ else:
101
+ self.meta['sid_num'] = self.meta.sid.map(sid_nums)
102
+
103
+ ix = np.random.choice(len(self), min(50000, len(self)), replace=False)
104
+ subset = self.patches[ix]
105
+ self.means = subset.mean(axis=(0,1,2))
106
+ self.stds = subset.std(axis=(0,1,2))
107
+ self.percentiles = np.percentile(np.abs(subset), percentile_thresh, axis=(0,1,2))
108
+ self.vmin = (-self.means - self.percentiles)/self.stds
109
+ self.vmax = (-self.means + self.percentiles)/self.stds
110
+ print(f'means: {self.means}')
111
+ print(f'stds: {self.stds}')
112
+
113
+ if standardize:
114
+ self.patches = (self.patches - self.means[None,None,None,:]) / self.stds[None,None,None,:]
115
+
116
+ def __len__(self):
117
+ return len(self.meta)
118
+
119
+ def __getitem__(self, idx):
120
+ patches = self.patches[idx]
121
+ sid_nums = self.meta.sid_num.values[idx]
122
+ if self.dim_order == 'numpy':
123
+ return patches, sid_nums
124
+ else:
125
+ return self.transform(patches), torch.tensor(sid_nums)