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 ADDED
@@ -0,0 +1,10 @@
1
+ from .spatial_integrate import *
2
+ from .spatial_map import *
3
+ from .logger import *
4
+ from .dataset import *
5
+ from .loss import *
6
+ from .model import *
7
+ from .preprocess import *
8
+ from .train_model import *
9
+ from .train import *
10
+ from .utils import *
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)