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.
@@ -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