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 +2 -0
- vima/association.py +77 -0
- vima/data/__init__.py +3 -0
- vima/data/ingest.py +441 -0
- vima/data/patchcollection.py +125 -0
- vima/data/samples.py +100 -0
- vima/models/__init__.py +0 -0
- vima/models/resnet_vae.py +68 -0
- vima/models/resnetlight_advanced_decoder.py +213 -0
- vima/models/resnetlight_advanced_encoder.py +217 -0
- vima/models/resnetlight_simple_decoder.py +192 -0
- vima/models/resnetlight_simple_encoder.py +195 -0
- vima/models/simple_vae.py +47 -0
- vima/models/vae.py +38 -0
- vima/training.py +199 -0
- vima/vis.py +279 -0
- vima_spatial-0.1.0.dist-info/METADATA +40 -0
- vima_spatial-0.1.0.dist-info/RECORD +20 -0
- vima_spatial-0.1.0.dist-info/WHEEL +5 -0
- vima_spatial-0.1.0.dist-info/top_level.txt +1 -0
vima/__init__.py
ADDED
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
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)
|