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
|
@@ -0,0 +1,240 @@
|
|
|
1
|
+
from fusemap.model import Fuse_network
|
|
2
|
+
from fusemap.preprocess import *
|
|
3
|
+
from fusemap.dataset import *
|
|
4
|
+
from fusemap.loss import *
|
|
5
|
+
from fusemap.config import *
|
|
6
|
+
from fusemap.utils import *
|
|
7
|
+
from fusemap.train_model import *
|
|
8
|
+
from torch.optim.lr_scheduler import ReduceLROnPlateau
|
|
9
|
+
import torch.distributions as D
|
|
10
|
+
from pathlib import Path
|
|
11
|
+
import itertools
|
|
12
|
+
import dgl.dataloading as dgl_dataload
|
|
13
|
+
import random
|
|
14
|
+
import os
|
|
15
|
+
import anndata as ad
|
|
16
|
+
import torch
|
|
17
|
+
import numpy as np
|
|
18
|
+
from tqdm import tqdm
|
|
19
|
+
import scanpy as sc
|
|
20
|
+
import dgl
|
|
21
|
+
import logging
|
|
22
|
+
try:
|
|
23
|
+
import pickle5 as pickle
|
|
24
|
+
except ModuleNotFoundError:
|
|
25
|
+
import pickle
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def spatial_integrate(
|
|
29
|
+
X_input,
|
|
30
|
+
args,
|
|
31
|
+
kneighbor,
|
|
32
|
+
input_identity,
|
|
33
|
+
data_pth=None,
|
|
34
|
+
):
|
|
35
|
+
### preprocess
|
|
36
|
+
ModelType.data_pth = data_pth
|
|
37
|
+
ModelType.save_dir = args.output_save_dir
|
|
38
|
+
ModelType.kneighbor = kneighbor
|
|
39
|
+
ModelType.input_identity = input_identity
|
|
40
|
+
ModelType.snapshot_path = f"{ModelType.save_dir}/snapshot.pt"
|
|
41
|
+
Path(f"{ModelType.save_dir}").mkdir(parents=True, exist_ok=True)
|
|
42
|
+
Path(f"{ModelType.save_dir}/trained_model").mkdir(parents=True, exist_ok=True)
|
|
43
|
+
|
|
44
|
+
ModelType.n_atlas = len(X_input)
|
|
45
|
+
preprocess_raw(
|
|
46
|
+
X_input,
|
|
47
|
+
ModelType.kneighbor,
|
|
48
|
+
ModelType.input_identity,
|
|
49
|
+
ModelType.use_input.value,
|
|
50
|
+
ModelType.n_atlas,
|
|
51
|
+
ModelType.data_pth,
|
|
52
|
+
)
|
|
53
|
+
for i in range(ModelType.n_atlas):
|
|
54
|
+
X_input[i].var.index = [i.upper() for i in X_input[i].var.index]
|
|
55
|
+
adatas = X_input
|
|
56
|
+
|
|
57
|
+
ModelType.n_obs = [adatas[i].shape[0] for i in range(ModelType.n_atlas)]
|
|
58
|
+
ModelType.input_dim = [adatas[i].n_vars for i in range(ModelType.n_atlas)]
|
|
59
|
+
ModelType.var_name = [list(i.var.index) for i in adatas]
|
|
60
|
+
|
|
61
|
+
all_unique_genes = sorted(list(get_allunique_gene_names(*ModelType.var_name)))
|
|
62
|
+
logging.info(
|
|
63
|
+
f"\n\nnumber of genes in each section:{[len(i) for i in ModelType.var_name]}, Number of all genes: {len(all_unique_genes)}.\n"
|
|
64
|
+
)
|
|
65
|
+
|
|
66
|
+
### create model
|
|
67
|
+
model = Fuse_network(
|
|
68
|
+
ModelType.pca_dim.value,
|
|
69
|
+
ModelType.input_dim,
|
|
70
|
+
ModelType.hidden_dim.value,
|
|
71
|
+
ModelType.latent_dim.value,
|
|
72
|
+
ModelType.dropout_rate.value,
|
|
73
|
+
ModelType.var_name,
|
|
74
|
+
all_unique_genes,
|
|
75
|
+
ModelType.use_input.value,
|
|
76
|
+
ModelType.n_atlas,
|
|
77
|
+
ModelType.input_identity,
|
|
78
|
+
ModelType.n_obs,
|
|
79
|
+
ModelType.n_epochs.value,
|
|
80
|
+
use_llm_gene_embedding=args.use_llm_gene_embedding
|
|
81
|
+
)
|
|
82
|
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
83
|
+
model.to(device)
|
|
84
|
+
if args.use_llm_gene_embedding=='combine':
|
|
85
|
+
model.ground_truth_rel_matrix=model.ground_truth_rel_matrix.to(device)
|
|
86
|
+
ModelType.use_llm_gene_embedding=args.use_llm_gene_embedding
|
|
87
|
+
|
|
88
|
+
ModelType.epochs_run_pretrain = 0
|
|
89
|
+
ModelType.epochs_run_final = 0
|
|
90
|
+
if os.path.exists(ModelType.snapshot_path):
|
|
91
|
+
logging.info("\n\nLoading snapshot\n")
|
|
92
|
+
load_snapshot(model, ModelType.snapshot_path, device)
|
|
93
|
+
|
|
94
|
+
### construct graph and data
|
|
95
|
+
adj_all, g_all = construct_data(
|
|
96
|
+
ModelType.n_atlas, adatas, ModelType.input_identity, model
|
|
97
|
+
)
|
|
98
|
+
feature_all = [
|
|
99
|
+
get_feature_sparse(device, adata.obsm["spatial_input"]) for adata in adatas
|
|
100
|
+
]
|
|
101
|
+
spatial_dataset_list = [
|
|
102
|
+
CustomGraphDataset(i, j, ModelType.use_input) for i, j in zip(g_all, adatas)
|
|
103
|
+
]
|
|
104
|
+
spatial_dataloader = CustomGraphDataLoader(
|
|
105
|
+
spatial_dataset_list,
|
|
106
|
+
dgl_dataload.MultiLayerFullNeighborSampler(1),
|
|
107
|
+
ModelType.batch_size.value,
|
|
108
|
+
shuffle=True,
|
|
109
|
+
n_atlas=ModelType.n_atlas,
|
|
110
|
+
drop_last=False,
|
|
111
|
+
)
|
|
112
|
+
spatial_dataloader_test = CustomGraphDataLoader(
|
|
113
|
+
spatial_dataset_list,
|
|
114
|
+
dgl_dataload.MultiLayerFullNeighborSampler(1),
|
|
115
|
+
ModelType.batch_size.value,
|
|
116
|
+
shuffle=False,
|
|
117
|
+
n_atlas=ModelType.n_atlas,
|
|
118
|
+
drop_last=False,
|
|
119
|
+
)
|
|
120
|
+
train_mask, val_mask = construct_mask(
|
|
121
|
+
ModelType.n_atlas, spatial_dataset_list, g_all
|
|
122
|
+
)
|
|
123
|
+
|
|
124
|
+
### train
|
|
125
|
+
flagconfig = FlagConfig()
|
|
126
|
+
if os.path.exists(f"{ModelType.save_dir}/lambda_disc_single.pkl"):
|
|
127
|
+
with open(f"{ModelType.save_dir}/lambda_disc_single.pkl", "rb") as openfile:
|
|
128
|
+
flagconfig.lambda_disc_single = pickle.load(openfile)
|
|
129
|
+
|
|
130
|
+
if not os.path.exists(
|
|
131
|
+
f"{ModelType.save_dir}/trained_model/FuseMap_pretrain_model_final.pt"
|
|
132
|
+
):
|
|
133
|
+
logging.info(
|
|
134
|
+
"\n\n---------------------------------- Phase 1. Pretrain FuseMap model ----------------------------------\n"
|
|
135
|
+
)
|
|
136
|
+
pretrain_model(
|
|
137
|
+
model,
|
|
138
|
+
spatial_dataloader,
|
|
139
|
+
feature_all,
|
|
140
|
+
adj_all,
|
|
141
|
+
device,
|
|
142
|
+
train_mask,
|
|
143
|
+
val_mask,
|
|
144
|
+
flagconfig,
|
|
145
|
+
)
|
|
146
|
+
|
|
147
|
+
if not os.path.exists(
|
|
148
|
+
f"{ModelType.save_dir}/latent_embeddings_all_single_pretrain.pkl"
|
|
149
|
+
):
|
|
150
|
+
logging.info(
|
|
151
|
+
"\n\n---------------------------------- Phase 2. Evaluate pretrained FuseMap model ----------------------------------\n"
|
|
152
|
+
)
|
|
153
|
+
if os.path.exists(
|
|
154
|
+
f"{ModelType.save_dir}/trained_model/FuseMap_pretrain_model_final.pt"
|
|
155
|
+
):
|
|
156
|
+
read_model(
|
|
157
|
+
model,
|
|
158
|
+
spatial_dataloader_test,
|
|
159
|
+
g_all,
|
|
160
|
+
feature_all,
|
|
161
|
+
adj_all,
|
|
162
|
+
device,
|
|
163
|
+
ModelType,
|
|
164
|
+
mode="pretrain",
|
|
165
|
+
)
|
|
166
|
+
else:
|
|
167
|
+
raise ValueError("No pretrained model!")
|
|
168
|
+
|
|
169
|
+
if not os.path.exists(f"{ModelType.save_dir}/balance_weight_single.pkl"):
|
|
170
|
+
logging.info(
|
|
171
|
+
"\n\n---------------------------------- Phase 3. Estimate_balancing_weight ----------------------------------\n"
|
|
172
|
+
)
|
|
173
|
+
balance_weight(model, adatas, ModelType.save_dir, ModelType.n_atlas, device)
|
|
174
|
+
|
|
175
|
+
if not os.path.exists(
|
|
176
|
+
f"{ModelType.save_dir}/trained_model/FuseMap_final_model_final.pt"
|
|
177
|
+
):
|
|
178
|
+
model.load_state_dict(
|
|
179
|
+
torch.load(
|
|
180
|
+
f"{ModelType.save_dir}/trained_model/FuseMap_pretrain_model_final.pt"
|
|
181
|
+
)
|
|
182
|
+
)
|
|
183
|
+
logging.info(
|
|
184
|
+
"\n\n---------------------------------- Phase 4. Train final FuseMap model ----------------------------------\n"
|
|
185
|
+
)
|
|
186
|
+
train_model(
|
|
187
|
+
model,
|
|
188
|
+
spatial_dataloader,
|
|
189
|
+
feature_all,
|
|
190
|
+
adj_all,
|
|
191
|
+
device,
|
|
192
|
+
train_mask,
|
|
193
|
+
val_mask,
|
|
194
|
+
flagconfig,
|
|
195
|
+
)
|
|
196
|
+
|
|
197
|
+
if not os.path.exists(
|
|
198
|
+
f"{ModelType.save_dir}/latent_embeddings_all_single_final.pkl"
|
|
199
|
+
):
|
|
200
|
+
logging.info(
|
|
201
|
+
"\n\n---------------------------------- Phase 5. Evaluate final FuseMap model ----------------------------------\n"
|
|
202
|
+
)
|
|
203
|
+
if os.path.exists(
|
|
204
|
+
f"{ModelType.save_dir}/trained_model/FuseMap_final_model_final.pt"
|
|
205
|
+
):
|
|
206
|
+
read_model(
|
|
207
|
+
model,
|
|
208
|
+
spatial_dataloader_test,
|
|
209
|
+
g_all,
|
|
210
|
+
feature_all,
|
|
211
|
+
adj_all,
|
|
212
|
+
device,
|
|
213
|
+
ModelType,
|
|
214
|
+
mode="final",
|
|
215
|
+
)
|
|
216
|
+
else:
|
|
217
|
+
raise ValueError("No final model!")
|
|
218
|
+
|
|
219
|
+
logging.info(
|
|
220
|
+
"\n\n---------------------------------- Finish ----------------------------------\n"
|
|
221
|
+
)
|
|
222
|
+
|
|
223
|
+
### read out gene embedding
|
|
224
|
+
read_gene_embedding(
|
|
225
|
+
model,
|
|
226
|
+
all_unique_genes,
|
|
227
|
+
ModelType.save_dir,
|
|
228
|
+
ModelType.n_atlas,
|
|
229
|
+
ModelType.var_name,
|
|
230
|
+
)
|
|
231
|
+
|
|
232
|
+
### read out cell embedding
|
|
233
|
+
read_cell_embedding(
|
|
234
|
+
adatas,
|
|
235
|
+
ModelType.save_dir,
|
|
236
|
+
args.keep_celltype,
|
|
237
|
+
args.keep_tissueregion,
|
|
238
|
+
)
|
|
239
|
+
|
|
240
|
+
return
|
fusemap/spatial_map.py
ADDED
|
@@ -0,0 +1,236 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
from fusemap.model import Fuse_network
|
|
3
|
+
from fusemap.preprocess import *
|
|
4
|
+
from fusemap.dataset import *
|
|
5
|
+
from fusemap.loss import *
|
|
6
|
+
from fusemap.config import *
|
|
7
|
+
from fusemap.utils import *
|
|
8
|
+
from fusemap.train_model import *
|
|
9
|
+
from torch.optim.lr_scheduler import ReduceLROnPlateau
|
|
10
|
+
import torch.distributions as D
|
|
11
|
+
from pathlib import Path
|
|
12
|
+
import itertools
|
|
13
|
+
import dgl.dataloading as dgl_dataload
|
|
14
|
+
import os
|
|
15
|
+
import anndata as ad
|
|
16
|
+
import torch
|
|
17
|
+
import numpy as np
|
|
18
|
+
from tqdm import tqdm
|
|
19
|
+
import scanpy as sc
|
|
20
|
+
|
|
21
|
+
try:
|
|
22
|
+
import pickle5 as pickle
|
|
23
|
+
except ModuleNotFoundError:
|
|
24
|
+
import pickle
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def spatial_map(
|
|
28
|
+
molccf_path,
|
|
29
|
+
X_input,
|
|
30
|
+
args,
|
|
31
|
+
kneighbor,
|
|
32
|
+
input_identity,
|
|
33
|
+
data_pth=None,
|
|
34
|
+
):
|
|
35
|
+
### preprocess
|
|
36
|
+
ModelType.data_pth = data_pth
|
|
37
|
+
ModelType.save_dir = args.output_save_dir
|
|
38
|
+
ModelType.kneighbor = kneighbor
|
|
39
|
+
ModelType.input_identity = input_identity
|
|
40
|
+
ModelType.snapshot_path = f"{ModelType.save_dir}/snapshot.pt"
|
|
41
|
+
Path(f"{ModelType.save_dir}").mkdir(parents=True, exist_ok=True)
|
|
42
|
+
Path(f"{ModelType.save_dir}/trained_model").mkdir(parents=True, exist_ok=True)
|
|
43
|
+
|
|
44
|
+
ModelType.n_atlas = len(X_input)
|
|
45
|
+
preprocess_raw(
|
|
46
|
+
X_input,
|
|
47
|
+
ModelType.kneighbor,
|
|
48
|
+
ModelType.input_identity,
|
|
49
|
+
ModelType.use_input.value,
|
|
50
|
+
ModelType.n_atlas,
|
|
51
|
+
ModelType.data_pth,
|
|
52
|
+
)
|
|
53
|
+
for i in range(ModelType.n_atlas):
|
|
54
|
+
X_input[i].var.index = [i.upper() for i in X_input[i].var.index]
|
|
55
|
+
adatas = X_input
|
|
56
|
+
|
|
57
|
+
ModelType.n_obs = [adatas[i].shape[0] for i in range(ModelType.n_atlas)]
|
|
58
|
+
ModelType.input_dim = [adatas[i].n_vars for i in range(ModelType.n_atlas)]
|
|
59
|
+
ModelType.var_name = [list(i.var.index) for i in adatas]
|
|
60
|
+
|
|
61
|
+
all_unique_genes = sorted(list(get_allunique_gene_names(*ModelType.var_name)))
|
|
62
|
+
logging.info(
|
|
63
|
+
f"\n\nnumber of genes in each section:{[len(i) for i in ModelType.var_name]}, Number of all genes: {len(all_unique_genes)}\n"
|
|
64
|
+
)
|
|
65
|
+
|
|
66
|
+
### load pretrain model
|
|
67
|
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
68
|
+
|
|
69
|
+
(
|
|
70
|
+
TRAINED_MODEL,
|
|
71
|
+
TRAINED_X_NUM,
|
|
72
|
+
TRAINED_GENE_EMBED,
|
|
73
|
+
TRAINED_GENE_NAME,
|
|
74
|
+
) = load_ref_model(molccf_path, device)
|
|
75
|
+
|
|
76
|
+
### define new model
|
|
77
|
+
PRETRAINED_GENE = []
|
|
78
|
+
new_train_gene = []
|
|
79
|
+
for i in all_unique_genes:
|
|
80
|
+
if i not in TRAINED_GENE_NAME:
|
|
81
|
+
new_train_gene.append(i)
|
|
82
|
+
else:
|
|
83
|
+
PRETRAINED_GENE.append(i)
|
|
84
|
+
pretrain_index = [TRAINED_GENE_NAME.index(i) for i in PRETRAINED_GENE]
|
|
85
|
+
logging.info(
|
|
86
|
+
f"\n\npretrain gene number:{len(PRETRAINED_GENE)}, new gene number:{len(new_train_gene)}\n"
|
|
87
|
+
)
|
|
88
|
+
|
|
89
|
+
adapt_model = Fuse_network(
|
|
90
|
+
ModelType.pca_dim.value,
|
|
91
|
+
ModelType.input_dim,
|
|
92
|
+
ModelType.hidden_dim.value,
|
|
93
|
+
ModelType.latent_dim.value,
|
|
94
|
+
ModelType.dropout_rate.value,
|
|
95
|
+
ModelType.var_name,
|
|
96
|
+
all_unique_genes,
|
|
97
|
+
ModelType.use_input.value,
|
|
98
|
+
ModelType.harmonized_gene,
|
|
99
|
+
ModelType.n_atlas,
|
|
100
|
+
ModelType.input_identity,
|
|
101
|
+
ModelType.n_obs,
|
|
102
|
+
ModelType.n_epochs.value,
|
|
103
|
+
pretrain_model=True,
|
|
104
|
+
pretrain_n_atlas=TRAINED_X_NUM,
|
|
105
|
+
PRETRAINED_GENE=PRETRAINED_GENE,
|
|
106
|
+
new_train_gene=new_train_gene,
|
|
107
|
+
)
|
|
108
|
+
adapt_model.to(device)
|
|
109
|
+
|
|
110
|
+
ModelType.epochs_run_pretrain = 0
|
|
111
|
+
ModelType.epochs_run_final = 0
|
|
112
|
+
if os.path.exists(ModelType.snapshot_path):
|
|
113
|
+
logging.info("\n\nLoading snapshot\n")
|
|
114
|
+
load_snapshot(adapt_model, ModelType.snapshot_path, device)
|
|
115
|
+
|
|
116
|
+
### construct graph and data
|
|
117
|
+
adj_all, g_all = construct_data(
|
|
118
|
+
ModelType.n_atlas, adatas, ModelType.input_identity, adapt_model
|
|
119
|
+
)
|
|
120
|
+
feature_all = [
|
|
121
|
+
get_feature_sparse(device, adata.obsm["spatial_input"]) for adata in adatas
|
|
122
|
+
]
|
|
123
|
+
spatial_dataset_list = [
|
|
124
|
+
CustomGraphDataset(i, j, ModelType.use_input) for i, j in zip(g_all, adatas)
|
|
125
|
+
]
|
|
126
|
+
spatial_dataloader = CustomGraphDataLoader(
|
|
127
|
+
spatial_dataset_list,
|
|
128
|
+
dgl_dataload.MultiLayerFullNeighborSampler(1),
|
|
129
|
+
ModelType.batch_size.value,
|
|
130
|
+
shuffle=True,
|
|
131
|
+
n_atlas=ModelType.n_atlas,
|
|
132
|
+
drop_last=False,
|
|
133
|
+
)
|
|
134
|
+
spatial_dataloader_test = CustomGraphDataLoader(
|
|
135
|
+
spatial_dataset_list,
|
|
136
|
+
dgl_dataload.MultiLayerFullNeighborSampler(1),
|
|
137
|
+
ModelType.batch_size.value,
|
|
138
|
+
shuffle=False,
|
|
139
|
+
n_atlas=ModelType.n_atlas,
|
|
140
|
+
drop_last=False,
|
|
141
|
+
)
|
|
142
|
+
train_mask, val_mask = construct_mask(
|
|
143
|
+
ModelType.n_atlas, spatial_dataset_list, g_all
|
|
144
|
+
)
|
|
145
|
+
|
|
146
|
+
### train
|
|
147
|
+
flagconfig = FlagConfig()
|
|
148
|
+
if os.path.exists(f"{ModelType.save_dir}/lambda_disc_single.pkl"):
|
|
149
|
+
with open(f"{ModelType.save_dir}/lambda_disc_single.pkl", "rb") as openfile:
|
|
150
|
+
flagconfig.lambda_disc_single = pickle.load(openfile)
|
|
151
|
+
|
|
152
|
+
if not os.path.exists(
|
|
153
|
+
f"{ModelType.save_dir}/trained_model/FuseMap_map_model_final.pt"
|
|
154
|
+
):
|
|
155
|
+
logging.info(
|
|
156
|
+
"\n\n---------------------------------- Phase 1. Map FuseMap model ----------------------------------\n"
|
|
157
|
+
)
|
|
158
|
+
### transfer model weight
|
|
159
|
+
transfer_weight(TRAINED_MODEL, pretrain_index, adapt_model)
|
|
160
|
+
|
|
161
|
+
### load reference data
|
|
162
|
+
(dataloader_pretrain_single, dataloader_pretrain_spatial) = load_ref_data(
|
|
163
|
+
molccf_path,
|
|
164
|
+
TRAINED_X_NUM,
|
|
165
|
+
ModelType.batch_size.value,
|
|
166
|
+
USE_REFERENCE_PCT=ModelType.USE_REFERENCE_PCT.value
|
|
167
|
+
)
|
|
168
|
+
|
|
169
|
+
map_model(
|
|
170
|
+
adapt_model,
|
|
171
|
+
spatial_dataloader,
|
|
172
|
+
feature_all,
|
|
173
|
+
adj_all,
|
|
174
|
+
device,
|
|
175
|
+
train_mask,
|
|
176
|
+
val_mask,
|
|
177
|
+
molccf_path,
|
|
178
|
+
dataloader_pretrain_single,
|
|
179
|
+
dataloader_pretrain_spatial,
|
|
180
|
+
TRAINED_X_NUM,
|
|
181
|
+
flagconfig,
|
|
182
|
+
)
|
|
183
|
+
|
|
184
|
+
if not os.path.exists(
|
|
185
|
+
f"{ModelType.save_dir}/latent_embeddings_all_single_map.pkl"
|
|
186
|
+
):
|
|
187
|
+
logging.info(
|
|
188
|
+
"\n\n---------------------------------- Phase 2. Evaluate mapped FuseMap model ----------------------------------\n"
|
|
189
|
+
)
|
|
190
|
+
if os.path.exists(
|
|
191
|
+
f"{ModelType.save_dir}/trained_model/FuseMap_map_model_final.pt"
|
|
192
|
+
):
|
|
193
|
+
read_model(
|
|
194
|
+
adapt_model,
|
|
195
|
+
spatial_dataloader_test,
|
|
196
|
+
g_all,
|
|
197
|
+
feature_all,
|
|
198
|
+
adj_all,
|
|
199
|
+
device,
|
|
200
|
+
ModelType,
|
|
201
|
+
mode="map",
|
|
202
|
+
)
|
|
203
|
+
else:
|
|
204
|
+
raise ValueError("No mapped model!")
|
|
205
|
+
|
|
206
|
+
|
|
207
|
+
logging.info(
|
|
208
|
+
"\n\n---------------------------------- Finish ----------------------------------\n"
|
|
209
|
+
)
|
|
210
|
+
|
|
211
|
+
### read out gene embedding
|
|
212
|
+
read_gene_embedding_map(
|
|
213
|
+
adapt_model,
|
|
214
|
+
new_train_gene,
|
|
215
|
+
PRETRAINED_GENE,
|
|
216
|
+
ModelType.save_dir,
|
|
217
|
+
ModelType.n_atlas,
|
|
218
|
+
ModelType.var_name,
|
|
219
|
+
)
|
|
220
|
+
|
|
221
|
+
### read out cell embedding
|
|
222
|
+
read_cell_embedding(
|
|
223
|
+
adatas,
|
|
224
|
+
ModelType.save_dir,
|
|
225
|
+
args.keep_celltype,
|
|
226
|
+
args.keep_tissueregion,
|
|
227
|
+
use_key='map',
|
|
228
|
+
)
|
|
229
|
+
|
|
230
|
+
### transfer molCCF cell annotations
|
|
231
|
+
transfer_annotation(
|
|
232
|
+
adatas,
|
|
233
|
+
ModelType.save_dir,
|
|
234
|
+
molccf_path,
|
|
235
|
+
)
|
|
236
|
+
return
|
fusemap/train.py
ADDED
|
@@ -0,0 +1,224 @@
|
|
|
1
|
+
from fusemap.model import Fuse_network
|
|
2
|
+
from fusemap.preprocess import *
|
|
3
|
+
from fusemap.dataset import *
|
|
4
|
+
from fusemap.loss import *
|
|
5
|
+
from fusemap.config import *
|
|
6
|
+
from fusemap.utils import *
|
|
7
|
+
from fusemap.train_model import *
|
|
8
|
+
from torch.optim.lr_scheduler import ReduceLROnPlateau
|
|
9
|
+
import torch.distributions as D
|
|
10
|
+
from pathlib import Path
|
|
11
|
+
import itertools
|
|
12
|
+
import dgl.dataloading as dgl_dataload
|
|
13
|
+
import random
|
|
14
|
+
import os
|
|
15
|
+
import anndata as ad
|
|
16
|
+
import torch
|
|
17
|
+
import numpy as np
|
|
18
|
+
from tqdm import tqdm
|
|
19
|
+
import scanpy as sc
|
|
20
|
+
import dgl
|
|
21
|
+
import logging
|
|
22
|
+
try:
|
|
23
|
+
import pickle5 as pickle
|
|
24
|
+
except ModuleNotFoundError:
|
|
25
|
+
import pickle
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def train(X_input, save_dir, kneighbor, input_identity,
|
|
32
|
+
data_pth=None, preprocess_save=False):
|
|
33
|
+
# ModelType = parse_ModelType()
|
|
34
|
+
|
|
35
|
+
ModelType.preprocess_save = preprocess_save
|
|
36
|
+
ModelType.data_pth = data_pth
|
|
37
|
+
ModelType.save_dir = save_dir
|
|
38
|
+
ModelType.kneighbor = kneighbor
|
|
39
|
+
ModelType.input_identity = input_identity
|
|
40
|
+
|
|
41
|
+
### preprocess
|
|
42
|
+
ModelType.snapshot_path = f"{ModelType.save_dir}/snapshot.pt"
|
|
43
|
+
Path(f"{ModelType.save_dir}").mkdir(parents=True, exist_ok=True)
|
|
44
|
+
Path(f"{ModelType.save_dir}/trained_model").mkdir(parents=True, exist_ok=True)
|
|
45
|
+
|
|
46
|
+
ModelType.n_atlas = len(X_input)
|
|
47
|
+
if ModelType.preprocess_save == False:
|
|
48
|
+
preprocess_raw(
|
|
49
|
+
X_input,
|
|
50
|
+
ModelType.kneighbor,
|
|
51
|
+
ModelType.input_identity,
|
|
52
|
+
ModelType.use_input.value,
|
|
53
|
+
ModelType.n_atlas,
|
|
54
|
+
ModelType.data_pth,
|
|
55
|
+
)
|
|
56
|
+
for i in range(ModelType.n_atlas):
|
|
57
|
+
X_input[i].var.index = [i.upper() for i in X_input[i].var.index]
|
|
58
|
+
adatas = X_input
|
|
59
|
+
ModelType.n_obs = [adatas[i].shape[0] for i in range(ModelType.n_atlas)]
|
|
60
|
+
ModelType.input_dim = [adatas[i].n_vars for i in range(ModelType.n_atlas)]
|
|
61
|
+
ModelType.var_name = [list(i.var.index) for i in adatas]
|
|
62
|
+
|
|
63
|
+
all_unique_genes = sorted(list(get_allunique_gene_names(*ModelType.var_name)))
|
|
64
|
+
logging.info(
|
|
65
|
+
f"\n\nnumber of genes in each section:{[len(i) for i in ModelType.var_name]}, Number of all genes: {len(all_unique_genes)}\n"
|
|
66
|
+
)
|
|
67
|
+
|
|
68
|
+
### model
|
|
69
|
+
model = Fuse_network(
|
|
70
|
+
ModelType.pca_dim.value,
|
|
71
|
+
ModelType.input_dim,
|
|
72
|
+
ModelType.hidden_dim.value,
|
|
73
|
+
ModelType.latent_dim.value,
|
|
74
|
+
ModelType.dropout_rate.value,
|
|
75
|
+
ModelType.var_name,
|
|
76
|
+
all_unique_genes,
|
|
77
|
+
ModelType.use_input.value,
|
|
78
|
+
ModelType.harmonized_gene,
|
|
79
|
+
ModelType.n_atlas,
|
|
80
|
+
ModelType.input_identity,
|
|
81
|
+
ModelType.n_obs,
|
|
82
|
+
ModelType.n_epochs.value,
|
|
83
|
+
)
|
|
84
|
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
85
|
+
model.to(device)
|
|
86
|
+
|
|
87
|
+
ModelType.epochs_run_pretrain = 0
|
|
88
|
+
ModelType.epochs_run_final = 0
|
|
89
|
+
if os.path.exists(ModelType.snapshot_path):
|
|
90
|
+
logging.info("\n\nLoading snapshot\n")
|
|
91
|
+
load_snapshot(model, ModelType.snapshot_path, device)
|
|
92
|
+
|
|
93
|
+
### construct graph and data
|
|
94
|
+
adj_all, g_all = construct_data(ModelType.n_atlas, adatas, ModelType.input_identity, model)
|
|
95
|
+
feature_all = [
|
|
96
|
+
get_feature_sparse(device, adata.obsm["spatial_input"]) for adata in adatas
|
|
97
|
+
]
|
|
98
|
+
spatial_dataset_list = [
|
|
99
|
+
CustomGraphDataset(i, j, ModelType.use_input) for i, j in zip(g_all, adatas)
|
|
100
|
+
]
|
|
101
|
+
spatial_dataloader = CustomGraphDataLoader(
|
|
102
|
+
spatial_dataset_list,
|
|
103
|
+
dgl_dataload.MultiLayerFullNeighborSampler(1),
|
|
104
|
+
ModelType.batch_size.value,
|
|
105
|
+
shuffle=True,
|
|
106
|
+
n_atlas=ModelType.n_atlas,
|
|
107
|
+
drop_last=False,
|
|
108
|
+
)
|
|
109
|
+
spatial_dataloader_test = CustomGraphDataLoader(
|
|
110
|
+
spatial_dataset_list,
|
|
111
|
+
dgl_dataload.MultiLayerFullNeighborSampler(1),
|
|
112
|
+
ModelType.batch_size.value,
|
|
113
|
+
shuffle=False,
|
|
114
|
+
n_atlas=ModelType.n_atlas,
|
|
115
|
+
drop_last=False,
|
|
116
|
+
)
|
|
117
|
+
train_mask, val_mask = construct_mask(ModelType.n_atlas, spatial_dataset_list, g_all)
|
|
118
|
+
|
|
119
|
+
### train
|
|
120
|
+
flagconfig=FlagConfig()
|
|
121
|
+
if os.path.exists(f"{ModelType.save_dir}/lambda_disc_single.pkl"):
|
|
122
|
+
with open(f"{ModelType.save_dir}/lambda_disc_single.pkl", "rb") as openfile:
|
|
123
|
+
flagconfig.lambda_disc_single = pickle.load(openfile)
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
if not os.path.exists(
|
|
127
|
+
f"{ModelType.save_dir}/trained_model/FuseMap_pretrain_model_final.pt"
|
|
128
|
+
):
|
|
129
|
+
logging.info(
|
|
130
|
+
"\n\n---------------------------------- Phase 1. Pretrain FuseMap model ----------------------------------\n"
|
|
131
|
+
)
|
|
132
|
+
pretrain_model(
|
|
133
|
+
model,
|
|
134
|
+
spatial_dataloader,
|
|
135
|
+
feature_all,
|
|
136
|
+
adj_all,
|
|
137
|
+
device,
|
|
138
|
+
train_mask,
|
|
139
|
+
val_mask,
|
|
140
|
+
flagconfig,
|
|
141
|
+
)
|
|
142
|
+
|
|
143
|
+
if not os.path.exists(f"{ModelType.save_dir}/latent_embeddings_all_single_pretrain.pkl"):
|
|
144
|
+
logging.info(
|
|
145
|
+
"\n\n---------------------------------- Phase 2. Evaluate pretrained FuseMap model ----------------------------------\n"
|
|
146
|
+
)
|
|
147
|
+
if os.path.exists(
|
|
148
|
+
f"{ModelType.save_dir}/trained_model/FuseMap_pretrain_model_final.pt"
|
|
149
|
+
):
|
|
150
|
+
read_model(
|
|
151
|
+
model,
|
|
152
|
+
spatial_dataloader_test,
|
|
153
|
+
g_all,
|
|
154
|
+
feature_all,
|
|
155
|
+
adj_all,
|
|
156
|
+
device,
|
|
157
|
+
ModelType,
|
|
158
|
+
mode="pretrain",
|
|
159
|
+
)
|
|
160
|
+
else:
|
|
161
|
+
raise ValueError("No pretrained model!")
|
|
162
|
+
|
|
163
|
+
if not os.path.exists(f"{ModelType.save_dir}/balance_weight_single.pkl"):
|
|
164
|
+
logging.info(
|
|
165
|
+
"\n\n---------------------------------- Phase 3. Estimate_balancing_weight ----------------------------------\n"
|
|
166
|
+
)
|
|
167
|
+
balance_weight(model, adatas, ModelType.save_dir, ModelType.n_atlas, device)
|
|
168
|
+
|
|
169
|
+
if not os.path.exists(
|
|
170
|
+
f"{ModelType.save_dir}/trained_model/FuseMap_final_model_final.pt"
|
|
171
|
+
):
|
|
172
|
+
model.load_state_dict(
|
|
173
|
+
torch.load(f"{ModelType.save_dir}/trained_model/FuseMap_pretrain_model_final.pt")
|
|
174
|
+
)
|
|
175
|
+
logging.info(
|
|
176
|
+
"\n\n---------------------------------- Phase 4. Train final FuseMap model ----------------------------------\n"
|
|
177
|
+
)
|
|
178
|
+
train_model(
|
|
179
|
+
model,
|
|
180
|
+
spatial_dataloader,
|
|
181
|
+
feature_all,
|
|
182
|
+
adj_all,
|
|
183
|
+
device,
|
|
184
|
+
train_mask,
|
|
185
|
+
val_mask,
|
|
186
|
+
flagconfig,
|
|
187
|
+
)
|
|
188
|
+
|
|
189
|
+
if not os.path.exists(f"{ModelType.save_dir}/latent_embeddings_all_single_final.pkl"):
|
|
190
|
+
logging.info(
|
|
191
|
+
"\n\n---------------------------------- Phase 5. Evaluate final FuseMap model ----------------------------------\n"
|
|
192
|
+
)
|
|
193
|
+
if os.path.exists(
|
|
194
|
+
f"{ModelType.save_dir}/trained_model/FuseMap_final_model_final.pt"
|
|
195
|
+
):
|
|
196
|
+
read_model(
|
|
197
|
+
model,
|
|
198
|
+
spatial_dataloader_test,
|
|
199
|
+
g_all,
|
|
200
|
+
feature_all,
|
|
201
|
+
adj_all,
|
|
202
|
+
device,
|
|
203
|
+
ModelType,
|
|
204
|
+
mode="final",
|
|
205
|
+
)
|
|
206
|
+
else:
|
|
207
|
+
raise ValueError("No final model!")
|
|
208
|
+
|
|
209
|
+
logging.info(
|
|
210
|
+
"\n\n---------------------------------- Finish ----------------------------------\n"
|
|
211
|
+
)
|
|
212
|
+
|
|
213
|
+
### read out gene embedding
|
|
214
|
+
read_gene_embedding(
|
|
215
|
+
model, all_unique_genes, ModelType.save_dir, ModelType.n_atlas, ModelType.var_name
|
|
216
|
+
)
|
|
217
|
+
|
|
218
|
+
### read out cell embedding
|
|
219
|
+
annotation_transfer(
|
|
220
|
+
adatas,
|
|
221
|
+
ModelType.save_dir,
|
|
222
|
+
)
|
|
223
|
+
|
|
224
|
+
return
|