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/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