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.
- vima_spatial-0.1.0/PKG-INFO +40 -0
- vima_spatial-0.1.0/README.md +8 -0
- vima_spatial-0.1.0/pyproject.toml +6 -0
- vima_spatial-0.1.0/setup.cfg +48 -0
- vima_spatial-0.1.0/src/vima/__init__.py +2 -0
- vima_spatial-0.1.0/src/vima/association.py +77 -0
- vima_spatial-0.1.0/src/vima/data/__init__.py +3 -0
- vima_spatial-0.1.0/src/vima/data/ingest.py +441 -0
- vima_spatial-0.1.0/src/vima/data/patchcollection.py +125 -0
- vima_spatial-0.1.0/src/vima/data/samples.py +100 -0
- vima_spatial-0.1.0/src/vima/models/__init__.py +0 -0
- vima_spatial-0.1.0/src/vima/models/resnet_vae.py +68 -0
- vima_spatial-0.1.0/src/vima/models/resnetlight_advanced_decoder.py +213 -0
- vima_spatial-0.1.0/src/vima/models/resnetlight_advanced_encoder.py +217 -0
- vima_spatial-0.1.0/src/vima/models/resnetlight_simple_decoder.py +192 -0
- vima_spatial-0.1.0/src/vima/models/resnetlight_simple_encoder.py +195 -0
- vima_spatial-0.1.0/src/vima/models/simple_vae.py +47 -0
- vima_spatial-0.1.0/src/vima/models/vae.py +38 -0
- vima_spatial-0.1.0/src/vima/training.py +199 -0
- vima_spatial-0.1.0/src/vima/vis.py +279 -0
- vima_spatial-0.1.0/src/vima_spatial.egg-info/PKG-INFO +40 -0
- vima_spatial-0.1.0/src/vima_spatial.egg-info/SOURCES.txt +24 -0
- vima_spatial-0.1.0/src/vima_spatial.egg-info/dependency_links.txt +1 -0
- vima_spatial-0.1.0/src/vima_spatial.egg-info/requires.txt +17 -0
- vima_spatial-0.1.0/src/vima_spatial.egg-info/top_level.txt +1 -0
|
@@ -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,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,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,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()
|