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.
- {fusemap-0.0.0 → fusemap-0.0.2}/LICENSE +0 -0
- {fusemap-0.0.0 → fusemap-0.0.2}/PKG-INFO +29 -9
- {fusemap-0.0.0 → fusemap-0.0.2}/README.md +2 -7
- fusemap-0.0.2/fusemap/__init__.py +10 -0
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/config.py +6 -4
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/dataset.py +28 -2
- fusemap-0.0.2/fusemap/deconvolution.py +375 -0
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/logger.py +0 -0
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/loss.py +399 -5
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/model.py +286 -9
- fusemap-0.0.2/fusemap/permutation.py +120 -0
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/preprocess.py +89 -29
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/spatial_integrate.py +31 -9
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/spatial_map.py +46 -17
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/train.py +32 -9
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/train_model.py +56 -33
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap/utils.py +157 -2
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap.egg-info/PKG-INFO +29 -9
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap.egg-info/SOURCES.txt +2 -0
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap.egg-info/dependency_links.txt +0 -0
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap.egg-info/entry_points.txt +0 -0
- fusemap-0.0.2/fusemap.egg-info/requires.txt +21 -0
- {fusemap-0.0.0 → fusemap-0.0.2}/fusemap.egg-info/top_level.txt +0 -0
- {fusemap-0.0.0 → fusemap-0.0.2}/setup.cfg +0 -0
- fusemap-0.0.2/setup.py +42 -0
- fusemap-0.0.0/fusemap/__init__.py +0 -10
- fusemap-0.0.0/fusemap.egg-info/requires.txt +0 -6
- fusemap-0.0.0/setup.py +0 -26
|
File without changes
|
|
@@ -1,8 +1,33 @@
|
|
|
1
|
-
Metadata-Version: 2.
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
2
|
Name: fusemap
|
|
3
|
-
Version: 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
|
|
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
|
|
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=
|
|
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
|
|
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
|