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/__init__.py
ADDED
fusemap/config.py
ADDED
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
import sys
|
|
2
|
+
import argparse
|
|
3
|
+
from typing import Dict, Any
|
|
4
|
+
from enum import Enum
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def parse_input_args():
|
|
8
|
+
parser = argparse.ArgumentParser(description="FuseMap")
|
|
9
|
+
|
|
10
|
+
parser.add_argument(
|
|
11
|
+
"--input_data_folder_path",
|
|
12
|
+
type=str,
|
|
13
|
+
required=True,
|
|
14
|
+
)
|
|
15
|
+
parser.add_argument(
|
|
16
|
+
"--output_save_dir",
|
|
17
|
+
type=str,
|
|
18
|
+
required=True,
|
|
19
|
+
)
|
|
20
|
+
parser.add_argument(
|
|
21
|
+
"--mode",
|
|
22
|
+
type=str,
|
|
23
|
+
required=True,
|
|
24
|
+
)
|
|
25
|
+
parser.add_argument(
|
|
26
|
+
"--keep_celltype",
|
|
27
|
+
type=str,
|
|
28
|
+
default="",
|
|
29
|
+
)
|
|
30
|
+
parser.add_argument(
|
|
31
|
+
"--keep_tissueregion",
|
|
32
|
+
type=str,
|
|
33
|
+
default="",
|
|
34
|
+
)
|
|
35
|
+
parser.add_argument(
|
|
36
|
+
"--use_llm_gene_embedding",
|
|
37
|
+
default=False,
|
|
38
|
+
)
|
|
39
|
+
|
|
40
|
+
args = parser.parse_args()
|
|
41
|
+
return args
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
class FlagConfig:
|
|
45
|
+
lambda_disc_single = 1
|
|
46
|
+
align_anneal = 1e10
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
class ModelType(Enum):
|
|
50
|
+
pca_dim = 50
|
|
51
|
+
hidden_dim = 512
|
|
52
|
+
latent_dim = 64
|
|
53
|
+
dropout_rate = 0.2
|
|
54
|
+
n_epochs = 16
|
|
55
|
+
batch_size = 64
|
|
56
|
+
learning_rate = 0.001
|
|
57
|
+
optim_kw = "RMSprop"
|
|
58
|
+
use_input = "norm"
|
|
59
|
+
lambda_ae_single = 1
|
|
60
|
+
lambda_disc_spatial = 1
|
|
61
|
+
lambda_ae_spatial = 1
|
|
62
|
+
align_noise_coef = 1.5
|
|
63
|
+
lr_patience_pretrain = 2
|
|
64
|
+
lr_factor_pretrain = 0.5
|
|
65
|
+
lr_limit_pretrain = 0.00001
|
|
66
|
+
patience_limit_final = 5
|
|
67
|
+
lr_patience_final = 3
|
|
68
|
+
lr_factor_final = 0.5
|
|
69
|
+
lr_limit_final = 0.00001
|
|
70
|
+
patience_limit_pretrain = 3
|
|
71
|
+
EPS = 1e-10
|
|
72
|
+
DIS_LAMDA = 2
|
|
73
|
+
TRAIN_WITHOUT_EVAL = 10
|
|
74
|
+
USE_REFERENCE_PCT = 0.02
|
|
75
|
+
verbose = False
|
fusemap/dataset.py
ADDED
|
@@ -0,0 +1,178 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
import scipy.sparse as sp
|
|
3
|
+
import dgl
|
|
4
|
+
import numpy as np
|
|
5
|
+
from torch.utils.data import Dataset, DataLoader, DataLoader, TensorDataset
|
|
6
|
+
import itertools
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def get_feature_sparse(device, feature):
|
|
10
|
+
return feature.copy() # .to(device)
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def construct_mask(n_atlas, spatial_dataset_list, g_all):
|
|
14
|
+
train_pct = 0.85
|
|
15
|
+
# np.random.seed(0)
|
|
16
|
+
num_train = [int(len(i) * train_pct) for i in spatial_dataset_list]
|
|
17
|
+
nodes_order = [np.random.permutation(g_i.number_of_nodes()) for g_i in g_all]
|
|
18
|
+
train_id = [
|
|
19
|
+
nodes_order_i[:num_train_i]
|
|
20
|
+
for nodes_order_i, num_train_i in zip(nodes_order, num_train)
|
|
21
|
+
]
|
|
22
|
+
# val_mask=[nodes_order_i[num_train_i:] for nodes_order_i,num_train_i in zip(nodes_order,num_train)]
|
|
23
|
+
train_mask = [
|
|
24
|
+
torch.zeros(
|
|
25
|
+
len(i),
|
|
26
|
+
)
|
|
27
|
+
for i in spatial_dataset_list
|
|
28
|
+
]
|
|
29
|
+
for i in range(n_atlas):
|
|
30
|
+
train_mask[i][train_id[i]] = 1
|
|
31
|
+
train_mask[i] = train_mask[i].bool()
|
|
32
|
+
val_mask = [~i for i in train_mask]
|
|
33
|
+
return train_mask, val_mask
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def construct_data(n_atlas, adatas, input_identity, model):
|
|
37
|
+
adj_all = []
|
|
38
|
+
g_all = []
|
|
39
|
+
for i in range(n_atlas):
|
|
40
|
+
adata = adatas[i]
|
|
41
|
+
if input_identity[i] == "ST":
|
|
42
|
+
adj_coo = adata.obsm["adj_normalized"].tocoo()
|
|
43
|
+
# adj_all.append(adj_coo.todense())
|
|
44
|
+
adj_all.append(adata.obsm["adj_normalized"])
|
|
45
|
+
else:
|
|
46
|
+
adj_raw = model.scrna_seq_adj["atlas" + str(i)]() # .weight
|
|
47
|
+
adj_coo = sp.coo_matrix(adj_raw.detach().cpu().numpy())
|
|
48
|
+
adj_all.append(adj_raw)
|
|
49
|
+
g_all.append(dgl.graph((adj_coo.row, adj_coo.col)))
|
|
50
|
+
return adj_all, g_all
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class CustomGraphDataset(Dataset):
|
|
54
|
+
def __init__(self, g, adata, useinput):
|
|
55
|
+
self.g = g
|
|
56
|
+
self.n_nodes = g.number_of_nodes()
|
|
57
|
+
|
|
58
|
+
def __len__(self):
|
|
59
|
+
return self.n_nodes
|
|
60
|
+
|
|
61
|
+
def __getitem__(self, idx):
|
|
62
|
+
# return X[idx], batch_idx[idx], library_size[idx], x_input[idx], idx
|
|
63
|
+
return idx
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
class CustomGraphDataLoader:
|
|
67
|
+
def __init__(self, dataset_all, sampler, batch_size, shuffle, n_atlas, drop_last):
|
|
68
|
+
self.dataset_all = dataset_all
|
|
69
|
+
self.sampler = sampler
|
|
70
|
+
self.batch_size = batch_size
|
|
71
|
+
self.shuffle = shuffle
|
|
72
|
+
self.n_atlas = n_atlas
|
|
73
|
+
|
|
74
|
+
self.dataloader = []
|
|
75
|
+
for i in range(n_atlas):
|
|
76
|
+
self.dataloader.append(
|
|
77
|
+
DataLoader(
|
|
78
|
+
self.dataset_all[i],
|
|
79
|
+
batch_size=batch_size,
|
|
80
|
+
shuffle=shuffle,
|
|
81
|
+
drop_last=drop_last,
|
|
82
|
+
)
|
|
83
|
+
)
|
|
84
|
+
cell_num = [len(i) for i in self.dataset_all]
|
|
85
|
+
self.max_value_index = np.argmax(cell_num)
|
|
86
|
+
|
|
87
|
+
def __iter__(self):
|
|
88
|
+
dataloader_iter_before = {}
|
|
89
|
+
dataloader_iter_after = {}
|
|
90
|
+
for i in np.arange(0, self.max_value_index):
|
|
91
|
+
dataloader_iter_before[i] = itertools.cycle(self.dataloader[i])
|
|
92
|
+
for i in np.arange(self.max_value_index + 1, self.n_atlas):
|
|
93
|
+
dataloader_iter_after[i] = itertools.cycle(self.dataloader[i])
|
|
94
|
+
|
|
95
|
+
for indices_max in self.dataloader[self.max_value_index]:
|
|
96
|
+
blocks = {}
|
|
97
|
+
for i in np.arange(0, self.max_value_index):
|
|
98
|
+
indices_i = next(dataloader_iter_before[i])
|
|
99
|
+
blocks[i] = {
|
|
100
|
+
"single": indices_i,
|
|
101
|
+
"spatial": self.sampler.sample_blocks(
|
|
102
|
+
self.dataset_all[i].g, indices_i
|
|
103
|
+
),
|
|
104
|
+
}
|
|
105
|
+
blocks[self.max_value_index] = {
|
|
106
|
+
"single": indices_max,
|
|
107
|
+
"spatial": self.sampler.sample_blocks(
|
|
108
|
+
self.dataset_all[self.max_value_index].g, indices_max
|
|
109
|
+
),
|
|
110
|
+
}
|
|
111
|
+
for i in np.arange(self.max_value_index + 1, self.n_atlas):
|
|
112
|
+
indices_i = next(dataloader_iter_after[i])
|
|
113
|
+
blocks[i] = {
|
|
114
|
+
"single": indices_i,
|
|
115
|
+
"spatial": self.sampler.sample_blocks(
|
|
116
|
+
self.dataset_all[i].g, indices_i
|
|
117
|
+
),
|
|
118
|
+
}
|
|
119
|
+
yield blocks
|
|
120
|
+
|
|
121
|
+
def __len__(self):
|
|
122
|
+
return max([len(i) for i in self.dataloader])
|
|
123
|
+
# return 100
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
class MapPretrainDataset(Dataset):
|
|
127
|
+
def __init__(self, X):
|
|
128
|
+
self.X = X
|
|
129
|
+
|
|
130
|
+
def __len__(self):
|
|
131
|
+
return len(self.X)
|
|
132
|
+
|
|
133
|
+
def __getitem__(self, idx):
|
|
134
|
+
return self.X[idx]
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
class MapPretrainDataLoader:
|
|
138
|
+
def __init__(self, dataset_all, batch_size, shuffle, n_atlas):
|
|
139
|
+
self.dataset_all = dataset_all
|
|
140
|
+
self.batch_size = batch_size
|
|
141
|
+
self.shuffle = shuffle
|
|
142
|
+
self.n_atlas = n_atlas
|
|
143
|
+
self.dataloader = []
|
|
144
|
+
for i in range(n_atlas):
|
|
145
|
+
self.dataloader.append(
|
|
146
|
+
DataLoader(
|
|
147
|
+
self.dataset_all[i],
|
|
148
|
+
batch_size=batch_size,
|
|
149
|
+
shuffle=shuffle,
|
|
150
|
+
drop_last=False,
|
|
151
|
+
)
|
|
152
|
+
)
|
|
153
|
+
cell_num = [len(i) for i in self.dataset_all]
|
|
154
|
+
self.max_value_index = np.argmax(cell_num)
|
|
155
|
+
|
|
156
|
+
def __iter__(self):
|
|
157
|
+
dataloader_iter_before = {}
|
|
158
|
+
dataloader_iter_after = {}
|
|
159
|
+
for i in np.arange(0, self.max_value_index):
|
|
160
|
+
dataloader_iter_before[i] = itertools.cycle(self.dataloader[i])
|
|
161
|
+
for i in np.arange(self.max_value_index + 1, self.n_atlas):
|
|
162
|
+
dataloader_iter_after[i] = itertools.cycle(self.dataloader[i])
|
|
163
|
+
|
|
164
|
+
for atlasdata_max in self.dataloader[self.max_value_index]:
|
|
165
|
+
blocks = {}
|
|
166
|
+
for i in np.arange(0, self.max_value_index):
|
|
167
|
+
atlasdata_i = next(dataloader_iter_before[i])
|
|
168
|
+
blocks[i] = atlasdata_i
|
|
169
|
+
|
|
170
|
+
blocks[self.max_value_index] = atlasdata_max
|
|
171
|
+
|
|
172
|
+
for i in np.arange(self.max_value_index + 1, self.n_atlas):
|
|
173
|
+
atlasdata_i = next(dataloader_iter_after[i])
|
|
174
|
+
blocks[i] = atlasdata_i
|
|
175
|
+
yield blocks
|
|
176
|
+
|
|
177
|
+
def __len__(self):
|
|
178
|
+
return max([len(i) for i in self.dataloader])
|
fusemap/logger.py
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
import os, re, logging
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
class MultipleHeaderFilter(logging.Filter):
|
|
5
|
+
def __init__(self, patterns_to_filter):
|
|
6
|
+
super().__init__()
|
|
7
|
+
self.patterns_to_filter = [re.compile(pattern) for pattern in patterns_to_filter]
|
|
8
|
+
|
|
9
|
+
def filter(self, record):
|
|
10
|
+
message = record.getMessage()
|
|
11
|
+
return not any(pattern.search(message) for pattern in self.patterns_to_filter)
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def setup_logging(save_path, patterns_to_filter=None):
|
|
15
|
+
"""
|
|
16
|
+
Configure logging to file and console, ignoring specified message patterns.
|
|
17
|
+
|
|
18
|
+
:param save_path: Path where the log file will be saved
|
|
19
|
+
:param patterns_to_filter: List of regex patterns for messages to be filtered out
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
if patterns_to_filter is None:
|
|
23
|
+
patterns_to_filter = [
|
|
24
|
+
r"^HTTP Request:",
|
|
25
|
+
r"^OpenAI API response:",
|
|
26
|
+
r"^Retrying request",
|
|
27
|
+
# Add more patterns here as needed
|
|
28
|
+
]
|
|
29
|
+
|
|
30
|
+
log_path = f"{save_path}/output.log"
|
|
31
|
+
if os.path.exists(log_path):
|
|
32
|
+
os.remove(log_path)
|
|
33
|
+
|
|
34
|
+
# Create handlers
|
|
35
|
+
file_handler = logging.FileHandler(log_path)
|
|
36
|
+
console_handler = logging.StreamHandler()
|
|
37
|
+
|
|
38
|
+
# Create formatter and add it to the handlers
|
|
39
|
+
formatter = logging.Formatter("%(asctime)s - %(levelname)s - %(message)s")
|
|
40
|
+
file_handler.setFormatter(formatter)
|
|
41
|
+
console_handler.setFormatter(formatter)
|
|
42
|
+
|
|
43
|
+
# Add filter to console handler not file handler
|
|
44
|
+
multiple_header_filter = MultipleHeaderFilter(patterns_to_filter)
|
|
45
|
+
# file_handler.addFilter(multiple_header_filter)
|
|
46
|
+
console_handler.addFilter(multiple_header_filter)
|
|
47
|
+
|
|
48
|
+
# Get the root logger and set its level
|
|
49
|
+
root_logger = logging.getLogger()
|
|
50
|
+
root_logger.setLevel(logging.INFO)
|
|
51
|
+
|
|
52
|
+
# Remove any existing handlers from the root logger
|
|
53
|
+
for handler in root_logger.handlers[:]:
|
|
54
|
+
root_logger.removeHandler(handler)
|
|
55
|
+
|
|
56
|
+
# Add the new handlers to the root logger
|
|
57
|
+
root_logger.addHandler(file_handler)
|
|
58
|
+
root_logger.addHandler(console_handler)
|