fusemap 0.0.0__tar.gz → 0.0.2__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.
Files changed (28) hide show
  1. {fusemap-0.0.0 → fusemap-0.0.2}/LICENSE +0 -0
  2. {fusemap-0.0.0 → fusemap-0.0.2}/PKG-INFO +29 -9
  3. {fusemap-0.0.0 → fusemap-0.0.2}/README.md +2 -7
  4. fusemap-0.0.2/fusemap/__init__.py +10 -0
  5. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/config.py +6 -4
  6. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/dataset.py +28 -2
  7. fusemap-0.0.2/fusemap/deconvolution.py +375 -0
  8. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/logger.py +0 -0
  9. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/loss.py +399 -5
  10. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/model.py +286 -9
  11. fusemap-0.0.2/fusemap/permutation.py +120 -0
  12. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/preprocess.py +89 -29
  13. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/spatial_integrate.py +31 -9
  14. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/spatial_map.py +46 -17
  15. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/train.py +32 -9
  16. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/train_model.py +56 -33
  17. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/utils.py +157 -2
  18. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap.egg-info/PKG-INFO +29 -9
  19. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap.egg-info/SOURCES.txt +2 -0
  20. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap.egg-info/dependency_links.txt +0 -0
  21. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap.egg-info/entry_points.txt +0 -0
  22. fusemap-0.0.2/fusemap.egg-info/requires.txt +21 -0
  23. {fusemap-0.0.0 → fusemap-0.0.2}/fusemap.egg-info/top_level.txt +0 -0
  24. {fusemap-0.0.0 → fusemap-0.0.2}/setup.cfg +0 -0
  25. fusemap-0.0.2/setup.py +42 -0
  26. fusemap-0.0.0/fusemap/__init__.py +0 -10
  27. fusemap-0.0.0/fusemap.egg-info/requires.txt +0 -6
  28. fusemap-0.0.0/setup.py +0 -26
File without changes
@@ -1,8 +1,33 @@
1
- Metadata-Version: 2.1
1
+ Metadata-Version: 2.4
2
2
  Name: fusemap
3
- Version: 0.0.0
3
+ Version: 0.0.2
4
4
  Description-Content-Type: text/markdown
5
5
  License-File: LICENSE
6
+ Requires-Dist: scanpy==1.9.3
7
+ Requires-Dist: torch==2.0.1
8
+ Requires-Dist: dgl==1.1.1
9
+ Requires-Dist: sparse==0.14.0
10
+ Requires-Dist: leidenalg==0.10.1
11
+ Requires-Dist: dglgo==0.0.2
12
+ Requires-Dist: streamlit==1.45.1
13
+ Requires-Dist: openai==1.82.1
14
+ Requires-Dist: langchain_anthropic==0.3.14
15
+ Requires-Dist: langchain==0.3.25
16
+ Requires-Dist: langchain-community==0.3.24
17
+ Requires-Dist: langchain-openai==0.3.18
18
+ Requires-Dist: audio-recorder-streamlit==0.0.10
19
+ Requires-Dist: plotly==6.1.2
20
+ Requires-Dist: harmonypy==0.0.10
21
+ Requires-Dist: langgraph==0.4.7
22
+ Requires-Dist: langgraph-prebuilt==0.2.2
23
+ Requires-Dist: langgraph-sdk==0.1.70
24
+ Requires-Dist: langsmith==0.3.43
25
+ Requires-Dist: easydict==1.13
26
+ Requires-Dist: numpy==1.26.2
27
+ Dynamic: description
28
+ Dynamic: description-content-type
29
+ Dynamic: license-file
30
+ Dynamic: requires-dist
6
31
 
7
32
  # FuseMap
8
33
  Integrate spatial transcripomics with universal gene, cell, and tissue embeddings.
@@ -44,14 +69,9 @@ scanpy
44
69
  seaborn
45
70
  ```
46
71
 
47
- ## Installation
48
- ```
49
- conda env create -f fusemap_environment.yml
50
- ```
51
-
52
72
 
53
- ## Tutorial
54
- - Read a [tutorial](./tutorial_sample_data.ipynb) on [sample data](https://drive.google.com/drive/folders/1DKfP5awTUa5gaL0WB-csD0M8v-COiBfY?usp=sharing) .
73
+ ## Installation and Tutorial
74
+ - Read the tutorial [here](https://fusemap.readthedocs.io/en/latest/) .
55
75
 
56
76
 
57
77
  ## Citation
@@ -38,14 +38,9 @@ scanpy
38
38
  seaborn
39
39
  ```
40
40
 
41
- ## Installation
42
- ```
43
- conda env create -f fusemap_environment.yml
44
- ```
45
-
46
41
 
47
- ## Tutorial
48
- - Read a [tutorial](./tutorial_sample_data.ipynb) on [sample data](https://drive.google.com/drive/folders/1DKfP5awTUa5gaL0WB-csD0M8v-COiBfY?usp=sharing) .
42
+ ## Installation and Tutorial
43
+ - Read the tutorial [here](https://fusemap.readthedocs.io/en/latest/) .
49
44
 
50
45
 
51
46
  ## Citation
@@ -0,0 +1,10 @@
1
+ from fusemap.spatial_integrate import *
2
+ from fusemap.spatial_map import *
3
+ from fusemap.logger import *
4
+ from fusemap.utils import *
5
+ from fusemap.config import *
6
+ from fusemap.model import *
7
+ from .dataset import *
8
+ from fusemap.loss import *
9
+ from fusemap.train_model import *
10
+ from fusemap.preprocess import *
@@ -1,6 +1,4 @@
1
- import sys
2
1
  import argparse
3
- from typing import Dict, Any
4
2
  from enum import Enum
5
3
 
6
4
 
@@ -34,9 +32,12 @@ def parse_input_args():
34
32
  )
35
33
  parser.add_argument(
36
34
  "--use_llm_gene_embedding",
37
- default=False,
35
+ default='false',
36
+ )
37
+ parser.add_argument(
38
+ "--pretrain_model_path",
39
+ default="",
38
40
  )
39
-
40
41
  args = parser.parse_args()
41
42
  return args
42
43
 
@@ -73,3 +74,4 @@ class ModelType(Enum):
73
74
  TRAIN_WITHOUT_EVAL = 10
74
75
  USE_REFERENCE_PCT = 0.02
75
76
  verbose = False
77
+ use_key='final'
@@ -2,15 +2,41 @@ import torch
2
2
  import scipy.sparse as sp
3
3
  import dgl
4
4
  import numpy as np
5
- from torch.utils.data import Dataset, DataLoader, DataLoader, TensorDataset
5
+ from torch.utils.data import Dataset, DataLoader, DataLoader
6
6
  import itertools
7
7
 
8
-
9
8
  def get_feature_sparse(device, feature):
10
9
  return feature.copy() # .to(device)
11
10
 
12
11
 
13
12
  def construct_mask(n_atlas, spatial_dataset_list, g_all):
13
+ """
14
+ Construct mask for training and validation
15
+
16
+ Parameters
17
+ ----------
18
+ n_atlas : int
19
+ Number of atlases
20
+ spatial_dataset_list : list
21
+ List of spatial datasets
22
+ g_all : list
23
+ List of graphs
24
+
25
+ Returns
26
+ -------
27
+ train_mask : list
28
+ List of training masks
29
+ val_mask : list
30
+ List of validation masks
31
+
32
+ Examples
33
+ --------
34
+ >>> n_atlas = 2
35
+ >>> spatial_dataset_list = [CustomGraphDataset(i, j, ModelType.use_input) for i, j in zip(g_all, adatas)]
36
+ >>> g_all = [dgl.graph((adj_coo.row, adj_coo.col)) for adj_coo in adj_all]
37
+ >>> train_mask, val_mask = construct_mask(n_atlas, spatial_dataset_list, g_all)
38
+
39
+ """
14
40
  train_pct = 0.85
15
41
  # np.random.seed(0)
16
42
  num_train = [int(len(i) * train_pct) for i in spatial_dataset_list]
@@ -0,0 +1,375 @@
1
+ # this script performs cell deconvolution based on Fusemap on starmap and stereomap
2
+ import scanpy as sc
3
+ import torch
4
+ from torch import optim
5
+ import numpy as np
6
+ import pandas as pd
7
+ from sklearn.cluster import KMeans
8
+ from anndata import AnnData
9
+ from tqdm import tqdm
10
+ from argparse import ArgumentParser
11
+ import json
12
+ import tangram as tg
13
+ from time import time
14
+
15
+ # Astrocytes: The five “Astr” types in list 2 were grouped to represent astrocytes.
16
+ # Dentate gyrus granule neurons: “GN DG” was chosen as the counterpart.
17
+ # Inhibitory and excitatory neurons: The various “EX…” and “IN…” subtypes in list 2 are grouped to cover several of the broad classes in list 1 (for example, “Di- and mesencephalon inhibitory neurons” and “Telencephalon inhibitory interneurons” are both mapped to subsets of “IN…” types).
18
+ # Missing matches: Some cell types from list 1 (for example, “Cerebellum neurons”, “Choroid epithelial cells”, “Enteric glia”, etc.) do not have clear counterparts in list 2 and are left with empty mappings.
19
+ # Extra types in list 2: For example, “Erythrocyte” in list 2 was not used because there was no matching entry in list 1.
20
+ cell_type_mapping_starmap_stereomap = {
21
+ "Astrocytes": ["Astr1", "Astr2", "Astr3", "Astr4", "Astr5"],
22
+ "Cerebellum neurons": [],
23
+ "Cholinergic and monoaminergic neurons": ["DA neuron"],
24
+ "Choroid epithelial cells": [],
25
+ "Dentate gyrus granule neurons": ["GN DG"],
26
+ "Dentate gyrus radial glia-like cells": [],
27
+ "Di- and mesencephalon excitatory neurons": ["EX", "EX Mb", "EX thalamus"],
28
+ "Di- and mesencephalon inhibitory neurons": [
29
+ "IN Pvalb+",
30
+ "IN Pvalb+Gad1+",
31
+ "IN Sst+",
32
+ "IN Vip+",
33
+ "IN thalamus",
34
+ ],
35
+ "Enteric glia": [],
36
+ "Ependymal cells": ["Ependymal"],
37
+ "Glutamatergic neuroblasts": [],
38
+ "Hindbrain neurons": [],
39
+ "Microglia": ["Microglia"],
40
+ "Non-glutamatergic neuroblasts": [],
41
+ "Olfactory ensheathing cells": [],
42
+ "Olfactory inhibitory neurons": [],
43
+ "Oligodendrocyte precursor cells": ["OPC"],
44
+ "Oligodendrocytes": ["Olig"],
45
+ "Peptidergic neurons": [],
46
+ "Pericytes": [],
47
+ "Perivascular macrophages": [],
48
+ "Spinal cord excitatory neurons": ["EX CA", "EX L2/3", "EX L4", "EX L5/6", "EX L6"],
49
+ "Spinal cord inhibitory neurons": [
50
+ "IN Pvalb+",
51
+ "IN Pvalb+Gad1+",
52
+ "IN Sst+",
53
+ "IN Vip+",
54
+ "IN thalamus",
55
+ ],
56
+ "Subcommissural organ hypendymal cells": [],
57
+ "Subventricular zone radial glia-like cells": [],
58
+ "Telencephalon inhibitory interneurons": [
59
+ "IN Pvalb+",
60
+ "IN Pvalb+Gad1+",
61
+ "IN Sst+",
62
+ "IN Vip+",
63
+ ],
64
+ "Telencephalon projecting excitatory neurons": [
65
+ "EX CA",
66
+ "EX L2/3",
67
+ "EX L4",
68
+ "EX L5/6",
69
+ "EX L6",
70
+ ],
71
+ "Telencephalon projecting inhibitory neurons": ["IN thalamus"],
72
+ "Vascular and leptomeningeal cells": ["Meninge"],
73
+ "Vascular endothelial cells": ["Endothelium"],
74
+ "Vascular smooth muscle cells": ["Smooth muscle cells"],
75
+ "nan": ["Unknown"],
76
+ }
77
+
78
+
79
+ def evaluate_spot_topk(M_final, spot_id: int, cell_type_mapping: dict, img_cate: list, ground_truth: list, k: int = 3):
80
+ # Get the top k predicted indices (sorted in descending order by score)
81
+ topk_indices = np.argsort(M_final[:, spot_id])[::-1][:k]
82
+ # Map these indices to their corresponding cell types
83
+ topk_predicted_types = [img_cate[i] for i in topk_indices]
84
+
85
+ # Combine the mappings for all top k predicted cell types
86
+ union_mapping = set()
87
+ for pred in topk_predicted_types:
88
+ # Use .get() to avoid KeyError if the prediction is not in the mapping
89
+ union_mapping.update(cell_type_mapping.get(pred, []))
90
+
91
+ # Get the ground truth cell type for the spot
92
+ cell_type_main = ground_truth[spot_id]
93
+
94
+ # Output the results
95
+ # print("Top-k predicted cell types:", topk_predicted_types)
96
+ # print("Union of corresponding ground truth mappings:", list(union_mapping))
97
+ # print("Ground truth cell type:", cell_type_main)
98
+ # print("Top-k correct:", cell_type_main in union_mapping)
99
+
100
+ if len(union_mapping) == 0 or cell_type_main == 'Unknown' or cell_type_main == 'Erythrocyte':
101
+ return None
102
+
103
+ return cell_type_main in union_mapping
104
+
105
+ def evaluate_topk_accuracy(M_final, cell_type_mapping, img_cate, ground_truth, n_spots, k: int = 1):
106
+ n_all, n_cor = 0, 0
107
+ for spot_id in tqdm(range(n_spots)):
108
+ res = evaluate_spot_topk(M_final, spot_id, cell_type_mapping, img_cate, ground_truth, k=k)
109
+ if res is None:
110
+ continue
111
+
112
+ n_all += 1
113
+ if res:
114
+ n_cor += 1
115
+
116
+ acc = n_cor / n_all
117
+ print(f"Top-{k} accuracy: {acc:.1%}, {n_cor}/{n_all}")
118
+
119
+ return acc
120
+ def get_cell_spot_embedding(ad_cell_embd: AnnData, cell_or_spot_column: str, cell_type_column: str = 'gtTaxonomyRank4'):
121
+ """
122
+ Extract embeddings and cell type information for cells and spots from an AnnData object.
123
+
124
+ Parameters:
125
+ - ad_cell_embd: AnnData
126
+ The AnnData object containing embeddings and metadata for cells and spots.
127
+ - cell_or_spot_column: str
128
+ The column name in `ad_cell_embd.obs` that indicates whether a row corresponds to a 'cell' or 'spot'.
129
+ - cell_type_column: str, optional (default: 'gtTaxonomyRank4')
130
+ The column name in `ad_cell_embd.obs` that contains cell type annotations for 'cell' rows.
131
+
132
+ Returns:
133
+ - cell_embd: np.array
134
+ The embedding matrix for rows corresponding to 'cell'.
135
+ - spot_embd: np.array
136
+ The embedding matrix for rows corresponding to 'spot'.
137
+ - cell_type: pd.Series
138
+ The cell type annotations for rows corresponding to 'cell'.
139
+ """
140
+ # Extract the embeddings for rows labeled as 'cell' in the specified column
141
+ cell_embd = ad_cell_embd.X[ad_cell_embd.obs[cell_or_spot_column] == 'cell']
142
+
143
+ # Extract the embeddings for rows labeled as 'spot' in the specified column
144
+ spot_embd = ad_cell_embd.X[ad_cell_embd.obs[cell_or_spot_column] == 'spot']
145
+
146
+ # Extract the cell type annotations for rows labeled as 'cell' in the specified column
147
+ cell_type = ad_cell_embd.obs[cell_type_column][ad_cell_embd.obs[cell_or_spot_column] == 'cell']
148
+
149
+ return cell_embd, spot_embd, cell_type
150
+
151
+ def get_representative_embeddings(Z_cells, cell_labels, n_types=None, n_prototypes=3, method="kmeans"):
152
+ """
153
+ Obtain representative embeddings (prototypes) for each cell type using clustering.
154
+
155
+ Parameters:
156
+ - Z_cells: np.array, shape [n_cells, C], the embedding matrix of all cells.
157
+ - cell_labels: np.array, shape [n_cells], the type labels for each cell.
158
+ - n_types: int, the total number of cell types.
159
+ - n_prototypes: int, the number of prototypes (representative embeddings) per cell type.
160
+ - method: str, clustering method, supports "kmeans".
161
+
162
+ Returns:
163
+ - Z_representative: np.array, shape [n_types, n_prototypes, C],
164
+ where each element Z_representative[k, :, :] corresponds to
165
+ the prototypes of cell type `k`.
166
+ """
167
+ C = Z_cells.shape[1]
168
+ n_types = np.max(cell_labels) + 1 if n_types is None else n_types
169
+
170
+ Z_representative = np.zeros((n_types, n_prototypes, C)) # Initialize the output array
171
+
172
+ for k in range(n_types):
173
+ # Get indices of cells belonging to the current cell type
174
+ indices = np.where(cell_labels == k)[0]
175
+ embeddings = Z_cells[indices, :] # Extract embeddings for this cell type
176
+
177
+ if embeddings.shape[0] == 0:
178
+ # If no cells belong to this type, fill with zeros
179
+ Z_representative[k, :, :] = np.zeros((n_prototypes, C))
180
+ continue
181
+
182
+ if method == "kmeans":
183
+ # Use K-Means clustering to find prototypes
184
+ kmeans = KMeans(n_clusters=n_prototypes, random_state=0, n_init=10)
185
+ kmeans.fit(embeddings)
186
+ Z_representative[k, :, :] = kmeans.cluster_centers_ # Shape: [n_prototypes, C]
187
+ else:
188
+ raise ValueError("Unsupported clustering method: choose 'kmeans'")
189
+
190
+ return Z_representative
191
+
192
+ # Define the sparse regularization term
193
+ def sparse_regularization(M):
194
+ # Count the number of non-zero elements along columns (L0 approximation)
195
+ # l0_norm = torch.sum(M != 0, dim=0)
196
+ # return torch.sum(l0_norm) / n_spots
197
+
198
+ return -torch.mean(M * torch.log(M + 1e-8))
199
+
200
+ def cosine_loss(pred, target):
201
+ # Normalize the vectors
202
+ pred_norm = pred / torch.norm(pred, dim=1, keepdim=True)
203
+ target_norm = target / torch.norm(target, dim=1, keepdim=True)
204
+ # Calculate cosine similarity for each row
205
+ cos_sim = torch.sum(pred_norm * target_norm, dim=1)
206
+ # Convert similarity to distance (1 - similarity) and take mean
207
+ return torch.mean(1 - cos_sim)
208
+
209
+ def optimize_cell_spot_assignment(Z_prototypes, Z_spots, lambda_reg=10, lr=0.03, num_epochs=2000, device='cpu'):
210
+ """
211
+ Optimize the cell-spot assignment matrix with sparse regularization.
212
+
213
+ Parameters:
214
+ - Z_prototypes: torch.Tensor
215
+ The embedding matrix of cell prototypes, shape (n_cells, embedding_dim).
216
+ - Z_spots: torch.Tensor
217
+ The embedding matrix of spots, shape (n_spots, embedding_dim).
218
+ - n_types: int
219
+ The number of cell types.
220
+ - n_prototypes: int
221
+ The number of prototypes per cell type.
222
+ - lambda_reg: float, optional (default: 10)
223
+ The regularization coefficient for sparsity.
224
+ - lr: float, optional (default: 0.03)
225
+ Learning rate for the optimizer.
226
+ - num_epochs: int, optional (default: 2000)
227
+ Number of optimization iterations.
228
+ - device: str, optional (default: 'cpu')
229
+ The device to use ('cpu' or 'cuda').
230
+
231
+ Returns:
232
+ - M_final: np.array
233
+ The final cell-spot assignment matrix, shape (n_types, n_spots).
234
+ """
235
+ # n_cells, n_spots = Z_prototypes.shape[0], Z_spots.shape[0]
236
+ n_types, n_prototypes, C = Z_prototypes.shape
237
+ n_spots = Z_spots.shape[0]
238
+
239
+ Z_prototypes = Z_prototypes.view(n_types * n_prototypes, C)
240
+
241
+ # Initialize the cell-spot assignment matrix M randomly
242
+ M_raw = torch.rand((n_types * n_prototypes, n_spots), requires_grad=True, device=device)
243
+
244
+ # Define the optimizer
245
+ optimizer = optim.Adam([M_raw], lr=lr)
246
+
247
+ # Define the loss function
248
+ # loss_fn = nn.MSELoss()
249
+
250
+ # Optimization loop
251
+ for epoch in tqdm(range(num_epochs)):
252
+ optimizer.zero_grad()
253
+
254
+ # Ensure that the columns of M sum to 1
255
+ M = torch.softmax(M_raw, dim=0)
256
+
257
+ # Reconstruction loss
258
+ reconstruction_loss = cosine_loss(torch.matmul(M.T, Z_prototypes), Z_spots)
259
+
260
+ # Sparse regularization loss
261
+ reg_loss = sparse_regularization(M)
262
+
263
+ # Total loss
264
+ total_loss = reconstruction_loss + lambda_reg * reg_loss
265
+
266
+ # Backpropagation
267
+ total_loss.backward()
268
+ optimizer.step()
269
+
270
+ # Print loss for monitoring
271
+ if epoch % 100 == 0:
272
+ print(f"Epoch {epoch}, Total Loss: {total_loss.item():.4f}, "
273
+ f"Reconstruction Loss: {reconstruction_loss.item():.4f}, "
274
+ f"Regularization Loss: {reg_loss.item():.4f}")
275
+
276
+ # Final optimized assignment matrix
277
+ M_opt = M.cpu().detach().numpy()
278
+ # n_types, n_prototypes, n_spots = M_opt.shape
279
+ # Aggregate by summing over prototypes for each cell type
280
+ M_final = M_opt.reshape(n_types, n_prototypes, n_spots).sum(1)
281
+
282
+ print("Optimization completed!")
283
+ return M_final
284
+
285
+ def get_args():
286
+ parser = ArgumentParser(description="Process number of prototypes and regularization parameter.")
287
+
288
+ parser.add_argument(
289
+ "--n_prototypes",
290
+ type=int,
291
+ default=5,
292
+ help="Number of prototypes (default: 5)"
293
+ )
294
+
295
+ parser.add_argument(
296
+ "--lambda_reg",
297
+ type=float,
298
+ default=0.1,
299
+ help="Regularization parameter lambda (default: 0.1)"
300
+ )
301
+
302
+ parser.add_argument(
303
+ "--n_epochs",
304
+ type=int,
305
+ default=1000,
306
+ )
307
+
308
+ parser.add_argument(
309
+ '--baseline',
310
+ action='store_true',
311
+ help='Use baseline method (implementing Tangram)'
312
+ )
313
+
314
+ args = parser.parse_args()
315
+ return args
316
+
317
+ if __name__ == '__main__':
318
+ # Load the cell embedding data
319
+ start_time = time()
320
+ args = get_args()
321
+ device = torch.device('mps')
322
+ torch.manual_seed(0)
323
+
324
+ ad_cell_embd = sc.read('/Users/mingzeyuan/Workspace/fusemap_deconvolution/raw_data/ad_celltype_embedding.h5ad')
325
+ ad_sc = sc.read_h5ad('/Users/mingzeyuan/Workspace/fusemap_deconvolution/raw_data/starmap.h5ad')
326
+ ad_sp = sc.read_h5ad('/Users/mingzeyuan/Workspace/fusemap_deconvolution/raw_data/stereoseq_mousebrain.h5ad')
327
+ # Create a new column 'cell_or_spot' with a default value (e.g., 'unknown')
328
+
329
+ if not args.baseline:
330
+ ad_cell_embd.obs['cell_or_spot'] = 'unknown'
331
+ # Assign values conditionally based on the 'batch' column
332
+ ad_cell_embd.obs.loc[ad_cell_embd.obs['batch'] == 'sample0', 'cell_or_spot'] = 'cell'
333
+ ad_cell_embd.obs.loc[ad_cell_embd.obs['batch'] == 'sample1', 'cell_or_spot'] = 'spot'
334
+
335
+ # Extract embeddings and cell type information
336
+ cell_embd, spot_embd, cell_type = get_cell_spot_embedding(ad_cell_embd, cell_or_spot_column='cell_or_spot', cell_type_column='gtTaxonomyRank4')
337
+ cell_label = pd.Categorical(cell_type).codes # Cell types for merscope
338
+
339
+ # Obtain representative embeddings for each cell type
340
+ Z_representative = get_representative_embeddings(cell_embd, cell_label, n_prototypes=args.n_prototypes, method="kmeans")
341
+
342
+ # Convert numpy arrays to torch tensors
343
+ Z_prototypes = torch.tensor(Z_representative, dtype=torch.float32, device=device)
344
+ Z_spots = torch.tensor(spot_embd, dtype=torch.float32, device=device)
345
+
346
+ else:
347
+ tg.pp_adatas(ad_sc, ad_sp, genes=None)
348
+ training_genes = ad_sc.uns['training_genes']
349
+ Z_cells = np.array(ad_sc[:, training_genes].X.toarray(), dtype="float32",)
350
+ Z_spots = np.array(ad_sp[:, training_genes].X.toarray(), dtype="float32",)
351
+ cell_types = pd.Categorical(ad_sc.obs['gtTaxonomyRank4']).codes # Cell types for merscope
352
+
353
+ Z_prototypes = get_representative_embeddings(Z_cells, cell_types, n_prototypes=args.n_prototypes).astype(np.float32)
354
+ Z_prototypes = torch.tensor(Z_prototypes).to(device)
355
+ Z_spots = torch.tensor(Z_spots).to(device)
356
+
357
+ # Optimize the cell-spot assignment matrix
358
+ M_final = optimize_cell_spot_assignment(Z_prototypes, Z_spots, lambda_reg=args.lambda_reg, lr=0.02, num_epochs=args.n_epochs, device=device)
359
+
360
+ # Save the final assignment matrix
361
+ suffix = f'{args.n_prototypes}_{args.lambda_reg}' if not args.baseline else f'baseline_{args.n_prototypes}_{args.lambda_reg}'
362
+ np.save(f'/Users/mingzeyuan/Workspace/fusemap_deconvolution/result/cell_spot_assignment_{suffix}.npy', M_final)
363
+
364
+ img_cate = pd.Categorical(ad_cell_embd.obs[ad_cell_embd.obs['batch'] == 'sample0']['gtTaxonomyRank4']).categories.tolist()
365
+
366
+ ground_truth = ad_sp.obs['gt_cell_type_main'].tolist()
367
+
368
+ res = {}
369
+ for k in [1, 2, 3, 4, 5]:
370
+ res[k] = evaluate_topk_accuracy(M_final, cell_type_mapping_starmap_stereomap, img_cate, ground_truth, ad_sp.X.shape[0], k=k)
371
+
372
+ res['time'] = time() - start_time
373
+
374
+ with open(f'/Users/mingzeyuan/Workspace/fusemap_deconvolution/result/accuracy_{suffix}.json', 'w') as f:
375
+ json.dump(res, f, indent=4)
File without changes