fusemap 0.0.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.
- fusemap/__init__.py +10 -0
- fusemap/config.py +75 -0
- fusemap/dataset.py +178 -0
- fusemap/logger.py +58 -0
- fusemap/loss.py +843 -0
- fusemap/model.py +362 -0
- fusemap/preprocess.py +161 -0
- fusemap/spatial_integrate.py +240 -0
- fusemap/spatial_map.py +236 -0
- fusemap/train.py +224 -0
- fusemap/train_model.py +1211 -0
- fusemap/utils.py +222 -0
- fusemap-0.0.0.dist-info/LICENSE +201 -0
- fusemap-0.0.0.dist-info/METADATA +67 -0
- fusemap-0.0.0.dist-info/RECORD +18 -0
- fusemap-0.0.0.dist-info/WHEEL +5 -0
- fusemap-0.0.0.dist-info/entry_points.txt +2 -0
- fusemap-0.0.0.dist-info/top_level.txt +1 -0
fusemap/model.py
ADDED
|
@@ -0,0 +1,362 @@
|
|
|
1
|
+
import torch.distributions as D
|
|
2
|
+
import torch.nn as nn
|
|
3
|
+
import torch.nn.functional as F
|
|
4
|
+
try:
|
|
5
|
+
import pickle5 as pickle
|
|
6
|
+
except ModuleNotFoundError:
|
|
7
|
+
import pickle
|
|
8
|
+
import torch
|
|
9
|
+
import numpy as np
|
|
10
|
+
import itertools
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def reset_parameters(para):
|
|
14
|
+
torch.nn.init.xavier_uniform_(para)
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class Discriminator(nn.Module):
|
|
18
|
+
def __init__(self, latent_dim, n_atlas, dropout_rate):
|
|
19
|
+
super(Discriminator, self).__init__()
|
|
20
|
+
|
|
21
|
+
self.linear_0 = nn.Linear(in_features=latent_dim, out_features=256, bias=True)
|
|
22
|
+
self.act_0 = nn.LeakyReLU(negative_slope=0.2)
|
|
23
|
+
self.dropout_0 = nn.Dropout(p=dropout_rate, inplace=False)
|
|
24
|
+
|
|
25
|
+
self.linear_1 = nn.Linear(in_features=256, out_features=256, bias=True)
|
|
26
|
+
self.act_1 = nn.LeakyReLU(negative_slope=0.2)
|
|
27
|
+
self.dropout_1 = nn.Dropout(p=dropout_rate, inplace=False)
|
|
28
|
+
|
|
29
|
+
self.pred = nn.Linear(in_features=256, out_features=n_atlas, bias=True)
|
|
30
|
+
# self.sigmoid = nn.Sigmoid()
|
|
31
|
+
|
|
32
|
+
def forward(self, x):
|
|
33
|
+
x = self.linear_0(x)
|
|
34
|
+
x = self.act_0(x)
|
|
35
|
+
x = self.dropout_0(x)
|
|
36
|
+
|
|
37
|
+
x = self.linear_1(x)
|
|
38
|
+
x = self.act_1(x)
|
|
39
|
+
x = self.dropout_1(x)
|
|
40
|
+
|
|
41
|
+
x = self.pred(x)
|
|
42
|
+
return x
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
class Adj_model(nn.Module):
|
|
46
|
+
def __init__(self, N):
|
|
47
|
+
super(Adj_model, self).__init__()
|
|
48
|
+
self.N = N
|
|
49
|
+
# initialize your weight
|
|
50
|
+
# self.weight = nn.Parameter(torch.full((N,N), 1.0/N))
|
|
51
|
+
self.weight = nn.Parameter(torch.empty((N, N)))
|
|
52
|
+
torch.nn.init.xavier_uniform_(self.weight)
|
|
53
|
+
|
|
54
|
+
def forward(self):
|
|
55
|
+
weight_relu = torch.relu(self.weight)
|
|
56
|
+
weight_relu = weight_relu + torch.eye(self.N).to(weight_relu.device)
|
|
57
|
+
|
|
58
|
+
### make weight matrix symmetric
|
|
59
|
+
weight_upper = torch.triu(weight_relu)
|
|
60
|
+
weight_lower = torch.tril(weight_relu).T
|
|
61
|
+
weight_symmetric = torch.max(weight_upper, weight_lower)
|
|
62
|
+
weight_symmetric = (
|
|
63
|
+
weight_symmetric
|
|
64
|
+
+ weight_symmetric.T
|
|
65
|
+
- torch.diag(weight_symmetric.diagonal())
|
|
66
|
+
)
|
|
67
|
+
|
|
68
|
+
k = 10 # torch.randint(3, 50, (1,)).item()
|
|
69
|
+
topk, _ = torch.topk(weight_symmetric, k, dim=1)
|
|
70
|
+
|
|
71
|
+
# Create a mask with 1s for the top k values and 0s for the rest
|
|
72
|
+
mask = weight_symmetric >= topk[:, -1:]
|
|
73
|
+
|
|
74
|
+
# Apply the mask to the weights
|
|
75
|
+
weight_topk = weight_symmetric * mask.float()
|
|
76
|
+
|
|
77
|
+
# Apply normalization along the row
|
|
78
|
+
weight_sum = (
|
|
79
|
+
torch.sum(weight_topk, dim=1, keepdim=True) + 1e-8
|
|
80
|
+
) # to prevent division by zero
|
|
81
|
+
weight_normalized = weight_topk / weight_sum
|
|
82
|
+
|
|
83
|
+
return weight_normalized
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
class FuseMapEncoder(nn.Module):
|
|
87
|
+
def __init__(
|
|
88
|
+
self, input_dim, hidden_dim, latent_dim, dropout_rate, normalization="batchnorm"
|
|
89
|
+
):
|
|
90
|
+
super(FuseMapEncoder, self).__init__()
|
|
91
|
+
self.dropout_0 = nn.Dropout(p=dropout_rate, inplace=False)
|
|
92
|
+
self.linear_0 = nn.Linear(input_dim, hidden_dim)
|
|
93
|
+
self.activation_0 = nn.LeakyReLU(negative_slope=0.2)
|
|
94
|
+
if normalization == "layernorm":
|
|
95
|
+
self.bn_0 = nn.LayerNorm(hidden_dim, eps=1e-05)
|
|
96
|
+
elif normalization == "batchnorm":
|
|
97
|
+
self.bn_0 = nn.BatchNorm1d(
|
|
98
|
+
hidden_dim,
|
|
99
|
+
eps=1e-05,
|
|
100
|
+
momentum=0.1,
|
|
101
|
+
affine=True,
|
|
102
|
+
track_running_stats=True,
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
self.dropout_1 = nn.Dropout(p=dropout_rate, inplace=False)
|
|
106
|
+
self.linear_1 = nn.Linear(hidden_dim, hidden_dim)
|
|
107
|
+
self.activation_1 = nn.LeakyReLU(negative_slope=0.2)
|
|
108
|
+
if normalization == "layernorm":
|
|
109
|
+
self.bn_1 = nn.LayerNorm(hidden_dim, eps=1e-05)
|
|
110
|
+
if normalization == "batchnorm":
|
|
111
|
+
self.bn_1 = nn.BatchNorm1d(
|
|
112
|
+
hidden_dim,
|
|
113
|
+
eps=1e-05,
|
|
114
|
+
momentum=0.1,
|
|
115
|
+
affine=True,
|
|
116
|
+
track_running_stats=True,
|
|
117
|
+
)
|
|
118
|
+
|
|
119
|
+
self.mean = nn.Linear(hidden_dim, latent_dim)
|
|
120
|
+
self.log_var = nn.Linear(hidden_dim, latent_dim)
|
|
121
|
+
|
|
122
|
+
def forward(self, x, adj):
|
|
123
|
+
h_1 = self.linear_0(x)
|
|
124
|
+
h_1 = self.bn_0(h_1)
|
|
125
|
+
h_1 = self.activation_0(h_1)
|
|
126
|
+
h_1 = self.dropout_0(h_1)
|
|
127
|
+
|
|
128
|
+
h_2 = self.linear_1(h_1)
|
|
129
|
+
h_2 = self.bn_1(h_2)
|
|
130
|
+
h_2 = self.activation_1(h_2)
|
|
131
|
+
h_2 = self.dropout_1(h_2)
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
z_mean = self.mean(h_2)
|
|
135
|
+
z_log_var = F.softplus(self.log_var(h_2))
|
|
136
|
+
|
|
137
|
+
z_sample = D.Normal(z_mean, z_log_var)
|
|
138
|
+
# z_sample_r = z_sample.rsample()
|
|
139
|
+
|
|
140
|
+
z_spatial = torch.mm(adj.T, z_mean)
|
|
141
|
+
|
|
142
|
+
return z_sample, None, z_spatial, z_mean
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
class FuseMapDecoder(nn.Module):
|
|
146
|
+
def __init__(self, gene_embedding, var_index):
|
|
147
|
+
super(FuseMapDecoder, self).__init__()
|
|
148
|
+
self.gene_embedding = gene_embedding
|
|
149
|
+
self.var_index = var_index
|
|
150
|
+
self.activation_3 = nn.LeakyReLU(negative_slope=0.2)
|
|
151
|
+
|
|
152
|
+
def forward(self, z_spatial, adj):
|
|
153
|
+
h_4 = torch.mm(adj, z_spatial)
|
|
154
|
+
x_recon_spatial = torch.mm(h_4, self.gene_embedding[:, self.var_index])
|
|
155
|
+
x_recon_spatial = self.activation_3(x_recon_spatial)
|
|
156
|
+
|
|
157
|
+
return x_recon_spatial
|
|
158
|
+
|
|
159
|
+
|
|
160
|
+
class FuseMapAdaptDecoder(nn.Module):
|
|
161
|
+
def __init__(self, var_index, gene_embedding_pretrain, gene_embedding_new):
|
|
162
|
+
super(FuseMapAdaptDecoder, self).__init__()
|
|
163
|
+
self.gene_embedding_pretrain = gene_embedding_pretrain
|
|
164
|
+
self.gene_embedding_new = gene_embedding_new
|
|
165
|
+
self.var_index = var_index
|
|
166
|
+
self.activation_3 = nn.LeakyReLU(negative_slope=0.2)
|
|
167
|
+
|
|
168
|
+
def forward(self, z, z_spatial, adj, gene_embedding_pretrain, gene_embedding_new):
|
|
169
|
+
h_4 = torch.mm(adj, z_spatial)
|
|
170
|
+
# p=0
|
|
171
|
+
# gene_embed_all = torch.hstack([gene_embedding_pretrain, self.gene_embedding_new ])
|
|
172
|
+
|
|
173
|
+
gene_embed_all = torch.hstack(
|
|
174
|
+
[
|
|
175
|
+
self.gene_embedding_new,
|
|
176
|
+
gene_embedding_pretrain,
|
|
177
|
+
]
|
|
178
|
+
)
|
|
179
|
+
|
|
180
|
+
x_recon_spatial = torch.mm(h_4, gene_embed_all[:, self.var_index])
|
|
181
|
+
x_recon_spatial = self.activation_3(x_recon_spatial)
|
|
182
|
+
|
|
183
|
+
return x_recon_spatial
|
|
184
|
+
|
|
185
|
+
|
|
186
|
+
class Fuse_network(nn.Module):
|
|
187
|
+
def __init__(
|
|
188
|
+
self,
|
|
189
|
+
pca_dim,
|
|
190
|
+
input_dim,
|
|
191
|
+
hidden_dim,
|
|
192
|
+
latent_dim,
|
|
193
|
+
dropout_rate,
|
|
194
|
+
var_name,
|
|
195
|
+
all_unique_genes,
|
|
196
|
+
use_input,
|
|
197
|
+
n_atlas,
|
|
198
|
+
input_identity,
|
|
199
|
+
n_obs,
|
|
200
|
+
num_epoch,
|
|
201
|
+
pretrain_model=False,
|
|
202
|
+
pretrain_n_atlas=0,
|
|
203
|
+
PRETRAINED_GENE=None,
|
|
204
|
+
new_train_gene=None,
|
|
205
|
+
use_llm_gene_embedding=False,
|
|
206
|
+
):
|
|
207
|
+
super(Fuse_network, self).__init__()
|
|
208
|
+
self.encoder = {}
|
|
209
|
+
self.decoder = {}
|
|
210
|
+
self.scrna_seq_adj = {}
|
|
211
|
+
|
|
212
|
+
if use_input == "norm" or use_input == "raw":
|
|
213
|
+
for i in range(n_atlas):
|
|
214
|
+
self.add_encoder_module(
|
|
215
|
+
"atlas" + str(i), input_dim[i], hidden_dim, latent_dim, dropout_rate
|
|
216
|
+
)
|
|
217
|
+
self.encoder = nn.ModuleDict(self.encoder)
|
|
218
|
+
|
|
219
|
+
##### build gene embedding
|
|
220
|
+
self.var_index = []
|
|
221
|
+
if use_llm_gene_embedding=='false':
|
|
222
|
+
if pretrain_model:
|
|
223
|
+
self.gene_embedding_pretrained = nn.Parameter(
|
|
224
|
+
torch.zeros(latent_dim, len(PRETRAINED_GENE))
|
|
225
|
+
)
|
|
226
|
+
self.gene_embedding_new = nn.Parameter(
|
|
227
|
+
torch.zeros(latent_dim, len(new_train_gene))
|
|
228
|
+
)
|
|
229
|
+
all_genes = new_train_gene + PRETRAINED_GENE
|
|
230
|
+
for ij in range(n_atlas):
|
|
231
|
+
self.var_index.append([all_genes.index(i) for i in var_name[ij]])
|
|
232
|
+
reset_parameters(self.gene_embedding_new)
|
|
233
|
+
else:
|
|
234
|
+
self.gene_embedding = nn.Parameter(
|
|
235
|
+
torch.zeros(latent_dim, len(all_unique_genes))
|
|
236
|
+
)
|
|
237
|
+
for ij in range(n_atlas):
|
|
238
|
+
self.var_index.append(
|
|
239
|
+
[all_unique_genes.index(i) for i in var_name[ij]]
|
|
240
|
+
)
|
|
241
|
+
reset_parameters(self.gene_embedding)
|
|
242
|
+
|
|
243
|
+
elif use_llm_gene_embedding=='combine':
|
|
244
|
+
if pretrain_model:
|
|
245
|
+
raise ValueError("pretrain_model is not supported for use_llm_gene_embedding")
|
|
246
|
+
else:
|
|
247
|
+
self.gene_embedding = nn.Parameter(
|
|
248
|
+
torch.zeros(latent_dim, len(all_unique_genes))
|
|
249
|
+
)
|
|
250
|
+
for ij in range(n_atlas):
|
|
251
|
+
self.var_index.append(
|
|
252
|
+
[all_unique_genes.index(i) for i in var_name[ij]]
|
|
253
|
+
)
|
|
254
|
+
reset_parameters(self.gene_embedding)
|
|
255
|
+
|
|
256
|
+
path_genept="./jupyter_notebook/data/GenePT_emebdding_v2/GenePT_gene_protein_embedding_model_3_text_pca.pickle"
|
|
257
|
+
with open(path_genept, "rb") as fp:
|
|
258
|
+
GPT_3_5_gene_embeddings = pickle.load(fp)
|
|
259
|
+
|
|
260
|
+
self.llm_gene_embedding = torch.zeros(latent_dim, len(all_unique_genes))
|
|
261
|
+
for i,gene in enumerate(all_unique_genes):
|
|
262
|
+
if gene in GPT_3_5_gene_embeddings.keys():
|
|
263
|
+
self.llm_gene_embedding[:,i] = torch.tensor(GPT_3_5_gene_embeddings[gene])
|
|
264
|
+
|
|
265
|
+
# Calculate gene embedding loss
|
|
266
|
+
ground_truth_matrix = self.llm_gene_embedding.T
|
|
267
|
+
ind = torch.sum(ground_truth_matrix,axis=1)!=0
|
|
268
|
+
ground_truth_matrix=ground_truth_matrix[ind,:]
|
|
269
|
+
|
|
270
|
+
self.llm_ind=ind
|
|
271
|
+
ground_truth_matrix_normalized = ground_truth_matrix / ground_truth_matrix.norm(dim=1, keepdim=True)
|
|
272
|
+
self.ground_truth_rel_matrix = torch.matmul(ground_truth_matrix_normalized, ground_truth_matrix_normalized.T)
|
|
273
|
+
|
|
274
|
+
elif use_llm_gene_embedding=='true':
|
|
275
|
+
if pretrain_model:
|
|
276
|
+
raise ValueError("pretrain_model is not supported for use_llm_gene_embedding")
|
|
277
|
+
else:
|
|
278
|
+
self.gene_embedding = torch.zeros(latent_dim, len(all_unique_genes))
|
|
279
|
+
for ij in range(n_atlas):
|
|
280
|
+
self.var_index.append(
|
|
281
|
+
[all_unique_genes.index(i) for i in var_name[ij]]
|
|
282
|
+
)
|
|
283
|
+
|
|
284
|
+
path_genept="./jupyter_notebook/data/GenePT_emebdding_v2/GenePT_gene_protein_embedding_model_3_text_pca.pickle"
|
|
285
|
+
with open(path_genept, "rb") as fp:
|
|
286
|
+
GPT_3_5_gene_embeddings = pickle.load(fp)
|
|
287
|
+
# reset_parameters(self.gene_embedding)
|
|
288
|
+
# ind=0
|
|
289
|
+
for i,gene in enumerate(all_unique_genes):
|
|
290
|
+
if gene in GPT_3_5_gene_embeddings.keys():
|
|
291
|
+
# print(gene)
|
|
292
|
+
# ind+=1
|
|
293
|
+
self.gene_embedding[:,i] = torch.tensor(GPT_3_5_gene_embeddings[gene])
|
|
294
|
+
self.gene_embedding=nn.Parameter(self.gene_embedding)
|
|
295
|
+
self.gene_embedding.requires_grad = False
|
|
296
|
+
else:
|
|
297
|
+
raise ValueError("use_llm_gene_embedding should be either 'true' or 'false' or 'combine'")
|
|
298
|
+
|
|
299
|
+
##### build decoders
|
|
300
|
+
if pretrain_model:
|
|
301
|
+
for ij in range(n_atlas):
|
|
302
|
+
self.add_adaptdecoder_module(
|
|
303
|
+
"atlas" + str(ij),
|
|
304
|
+
self.var_index[ij],
|
|
305
|
+
self.gene_embedding_pretrained,
|
|
306
|
+
self.gene_embedding_new,
|
|
307
|
+
)
|
|
308
|
+
else:
|
|
309
|
+
for ij in range(n_atlas):
|
|
310
|
+
self.add_decoder_module(
|
|
311
|
+
"atlas" + str(ij),
|
|
312
|
+
self.gene_embedding,
|
|
313
|
+
self.var_index[ij],
|
|
314
|
+
)
|
|
315
|
+
self.decoder = nn.ModuleDict(self.decoder)
|
|
316
|
+
|
|
317
|
+
##### build discriminators
|
|
318
|
+
self.discriminator_single = Discriminator(latent_dim, n_atlas, dropout_rate)
|
|
319
|
+
self.discriminator_spatial = Discriminator(latent_dim, n_atlas, dropout_rate)
|
|
320
|
+
|
|
321
|
+
if pretrain_model:
|
|
322
|
+
self.discriminator_single_pretrain = Discriminator(
|
|
323
|
+
latent_dim, pretrain_n_atlas, dropout_rate
|
|
324
|
+
)
|
|
325
|
+
self.discriminator_spatial_pretrain = Discriminator(
|
|
326
|
+
latent_dim, pretrain_n_atlas, dropout_rate
|
|
327
|
+
)
|
|
328
|
+
|
|
329
|
+
##### build scrnaseq adjacency matrix
|
|
330
|
+
for i in range(n_atlas):
|
|
331
|
+
if input_identity[i] == "scrna":
|
|
332
|
+
self.scrna_seq_adj["atlas" + str(i)] = Adj_model(n_obs[i])
|
|
333
|
+
self.scrna_seq_adj = nn.ModuleDict(self.scrna_seq_adj)
|
|
334
|
+
|
|
335
|
+
def add_encoder_module(
|
|
336
|
+
self, key, input_dim, hidden_dim, latent_dim, dropout_rate=0.1
|
|
337
|
+
):
|
|
338
|
+
self.encoder[key] = FuseMapEncoder(
|
|
339
|
+
input_dim, hidden_dim, latent_dim, dropout_rate
|
|
340
|
+
)
|
|
341
|
+
|
|
342
|
+
def add_decoder_module(self, key, gene_embedding, var_index):
|
|
343
|
+
self.decoder[key] = FuseMapDecoder(gene_embedding, var_index)
|
|
344
|
+
|
|
345
|
+
def add_adaptdecoder_module(self, key, var_index, gene_pretrain, gene_new):
|
|
346
|
+
self.decoder[key] = FuseMapAdaptDecoder(var_index, gene_pretrain, gene_new)
|
|
347
|
+
|
|
348
|
+
|
|
349
|
+
|
|
350
|
+
|
|
351
|
+
class NNTransfer(nn.Module):
|
|
352
|
+
def __init__(self, input_dim=128, output_dim=10):
|
|
353
|
+
super(NNTransfer, self).__init__()
|
|
354
|
+
self.fc1 = nn.Linear(input_dim, 256)
|
|
355
|
+
self.fc2 = nn.Linear(256, output_dim)
|
|
356
|
+
self.activate = nn.Softmax(dim=1)
|
|
357
|
+
|
|
358
|
+
def forward(self, x):
|
|
359
|
+
x = torch.relu(self.fc1(x))
|
|
360
|
+
x = self.fc2(x)
|
|
361
|
+
x= self.activate(x)
|
|
362
|
+
return x
|
fusemap/preprocess.py
ADDED
|
@@ -0,0 +1,161 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
import scipy.sparse as sp
|
|
3
|
+
import scipy
|
|
4
|
+
from scipy.spatial import Delaunay
|
|
5
|
+
from scipy.sparse.csr import csr_matrix
|
|
6
|
+
import scanpy as sc
|
|
7
|
+
from sklearn.neighbors import NearestNeighbors
|
|
8
|
+
import logging
|
|
9
|
+
|
|
10
|
+
def preprocess_raw(
|
|
11
|
+
X_input, kneighbor, input_identity, use_input, n_atlas, data_pth=None
|
|
12
|
+
):
|
|
13
|
+
logging.info(
|
|
14
|
+
"\n\n---------------------------------- Preprocess adata ----------------------------------\n"
|
|
15
|
+
)
|
|
16
|
+
X_input = preprocess_adata(X_input, n_atlas)
|
|
17
|
+
|
|
18
|
+
logging.info(
|
|
19
|
+
"\n\n---------------------------------- Construct graph adata ----------------------------------\n"
|
|
20
|
+
)
|
|
21
|
+
construct_graph(X_input, n_atlas, kneighbor, input_identity)
|
|
22
|
+
|
|
23
|
+
logging.info(
|
|
24
|
+
"\n\n---------------------------------- Process graph adata ----------------------------------\n"
|
|
25
|
+
)
|
|
26
|
+
preprocess_adj_sparse(X_input, n_atlas, input_identity)
|
|
27
|
+
|
|
28
|
+
get_spatial_input(X_input, n_atlas, use_input)
|
|
29
|
+
|
|
30
|
+
if data_pth is not None:
|
|
31
|
+
for pth_i, data_i in zip(data_pth, X_input[: len(data_pth)]):
|
|
32
|
+
logging.info(f"\n\nSaving processed data in {pth_i}\n")
|
|
33
|
+
data_i.write_h5ad(pth_i)
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def preprocess_adata(X_input, n_atlas):
|
|
37
|
+
for i in range(n_atlas):
|
|
38
|
+
if not "spatial_input" in X_input[i].obsm:
|
|
39
|
+
### filter genes
|
|
40
|
+
if isinstance(X_input[i].X, np.ndarray):
|
|
41
|
+
X_input[i] = X_input[i][:, np.sum(X_input[i].X, axis=0) > 5]
|
|
42
|
+
X_input[i] = X_input[i][:, np.max(X_input[i].X, axis=0) > 3]
|
|
43
|
+
if scipy.sparse.issparse(X_input[i].X):
|
|
44
|
+
X_input[i] = X_input[i][:, np.sum(X_input[i].X.toarray(), axis=0) > 5]
|
|
45
|
+
X_input[i] = X_input[i][:, np.max(X_input[i].X.toarray(), axis=0) > 3]
|
|
46
|
+
|
|
47
|
+
### unify genes
|
|
48
|
+
X_input[i].var.index = [i.upper() for i in X_input[i].var.index]
|
|
49
|
+
|
|
50
|
+
### keep unique genes
|
|
51
|
+
_, indices = np.unique(X_input[i].var.index, return_index=True)
|
|
52
|
+
X_input[i] = X_input[i][:, indices]
|
|
53
|
+
|
|
54
|
+
### unify genes
|
|
55
|
+
X_input[i].var.index = [i.upper() for i in X_input[i].var.index]
|
|
56
|
+
|
|
57
|
+
### filter cells
|
|
58
|
+
X_input[i] = X_input[i][np.sum(X_input[i].X, axis=1) > 5]
|
|
59
|
+
|
|
60
|
+
### normalize and pca
|
|
61
|
+
X_input[i].layers["counts"] = X_input[i].X.copy()
|
|
62
|
+
|
|
63
|
+
sc.pp.normalize_total(X_input[i]) # , target_sum=1e4)
|
|
64
|
+
sc.pp.log1p(X_input[i])
|
|
65
|
+
sc.pp.scale(X_input[i], zero_center=False, max_value=10)
|
|
66
|
+
if isinstance(X_input[i].X, np.ndarray):
|
|
67
|
+
X_input[i].X = csr_matrix(X_input[i].X)
|
|
68
|
+
|
|
69
|
+
return X_input
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def construct_graph(adatas, n_atlas, kneighbor, input_identity):
|
|
73
|
+
for i_atlas in range(n_atlas):
|
|
74
|
+
if input_identity[i_atlas] == "ST":
|
|
75
|
+
if not "adj_normalized" in adatas[i_atlas].obsm:
|
|
76
|
+
adata = adatas[i_atlas]
|
|
77
|
+
k = kneighbor[i_atlas]
|
|
78
|
+
data = np.array(adata.obs[["x", "y"]])
|
|
79
|
+
if k == "delaunay":
|
|
80
|
+
tri = Delaunay(data)
|
|
81
|
+
indptr, indices = tri.vertex_neighbor_vertices
|
|
82
|
+
adjacency_matrix = csr_matrix(
|
|
83
|
+
(np.ones_like(indices, dtype=np.float64), indices, indptr),
|
|
84
|
+
shape=(data.shape[0], data.shape[0]),
|
|
85
|
+
)
|
|
86
|
+
if k == "delaunay3d":
|
|
87
|
+
data = np.array(adata.obs[["x", "y", "z"]])
|
|
88
|
+
tri = Delaunay(data)
|
|
89
|
+
indptr, indices = tri.vertex_neighbor_vertices
|
|
90
|
+
adjacency_matrix = csr_matrix(
|
|
91
|
+
(np.ones_like(indices, dtype=np.float64), indices, indptr),
|
|
92
|
+
shape=(data.shape[0], data.shape[0]),
|
|
93
|
+
)
|
|
94
|
+
if "knn" in k:
|
|
95
|
+
if "3d" in k:
|
|
96
|
+
data = np.array(adata.obs[["x", "y", "z"]])
|
|
97
|
+
|
|
98
|
+
knn_k = 10
|
|
99
|
+
nbrs = NearestNeighbors(
|
|
100
|
+
n_neighbors=knn_k + 1, algorithm="auto"
|
|
101
|
+
).fit(data)
|
|
102
|
+
distances, indices = nbrs.kneighbors(data)
|
|
103
|
+
|
|
104
|
+
# Create an adjacency matrix
|
|
105
|
+
num_spots = data.shape[0]
|
|
106
|
+
adjacency_matrix = np.zeros((num_spots, num_spots))
|
|
107
|
+
|
|
108
|
+
for i in range(num_spots):
|
|
109
|
+
# indices[i, 1:] to exclude the point itself (the first nearest neighbor)
|
|
110
|
+
for j in indices[i, 1:]:
|
|
111
|
+
adjacency_matrix[i, j] = 1
|
|
112
|
+
adjacency_matrix[
|
|
113
|
+
j, i
|
|
114
|
+
] = 1 # Because it's an undirected graph
|
|
115
|
+
|
|
116
|
+
adata.obsm["adj"] = adjacency_matrix
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
def preprocess_adj_sparse(adatas, n_atlas, input_identity):
|
|
120
|
+
for i in range(n_atlas):
|
|
121
|
+
if input_identity[i] == "ST":
|
|
122
|
+
if not "adj_normalized" in adatas[i].obsm:
|
|
123
|
+
adata = adatas[i]
|
|
124
|
+
adj = sp.coo_matrix(adata.obsm["adj"])
|
|
125
|
+
adj_ = adj + sp.eye(adj.shape[0])
|
|
126
|
+
rowsum = np.array(adj_.sum(1))
|
|
127
|
+
degree_mat_inv_sqrt = sp.diags(np.power(rowsum, -0.5).flatten())
|
|
128
|
+
adj_normalized = (
|
|
129
|
+
adj_.dot(degree_mat_inv_sqrt)
|
|
130
|
+
.transpose()
|
|
131
|
+
.dot(degree_mat_inv_sqrt)
|
|
132
|
+
.tocoo()
|
|
133
|
+
)
|
|
134
|
+
adata.obsm[
|
|
135
|
+
"adj_normalized"
|
|
136
|
+
] = adj_normalized # sparse_mx_to_torch_sparse_tensor(adj_normalized)
|
|
137
|
+
adata.obsm["adj_normalized"] = adata.obsm["adj_normalized"].tocsr()
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
def get_spatial_input(adatas, n_atlas, use_input):
|
|
141
|
+
for i_atlas in range(n_atlas):
|
|
142
|
+
adata = adatas[i_atlas]
|
|
143
|
+
if not "spatial_input" in adata.obsm:
|
|
144
|
+
if use_input == "pca":
|
|
145
|
+
adata.obsm["spatial_input"] = adata.obsm["X_pca"]
|
|
146
|
+
if use_input == "raw":
|
|
147
|
+
adata.obsm["spatial_input"] = adata.layers["counts"]
|
|
148
|
+
if use_input == "norm":
|
|
149
|
+
adata.obsm["spatial_input"] = adata.X
|
|
150
|
+
|
|
151
|
+
|
|
152
|
+
def get_unique_gene_indices(gene_list):
|
|
153
|
+
unique_genes, indices = np.unique(gene_list, return_index=True)
|
|
154
|
+
return indices
|
|
155
|
+
|
|
156
|
+
|
|
157
|
+
def get_allunique_gene_names(*sample_gene_lists):
|
|
158
|
+
unique_genes = set()
|
|
159
|
+
for gene_list in sample_gene_lists:
|
|
160
|
+
unique_genes.update(gene_list)
|
|
161
|
+
return unique_genes
|