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/train_model.py ADDED
@@ -0,0 +1,1211 @@
1
+ import logging
2
+ from torch.optim.lr_scheduler import ReduceLROnPlateau
3
+ import torch.distributions as D
4
+ from pathlib import Path
5
+ import itertools
6
+ import dgl.dataloading as dgl_dataload
7
+ import random
8
+ import os
9
+ from fusemap.config import *
10
+ from fusemap.dataset import *
11
+ from fusemap.utils import *
12
+ from fusemap.loss import *
13
+ import anndata as ad
14
+ import torch
15
+ import numpy as np
16
+ from tqdm import tqdm
17
+ import scanpy as sc
18
+ import dgl
19
+ import torch.nn as nn
20
+
21
+ try:
22
+ import pickle5 as pickle
23
+ except ModuleNotFoundError:
24
+ import pickle
25
+
26
+
27
+ def get_data(blocks_all, feature_all, adj_all, train_mask, device, model):
28
+ row_index_all = {}
29
+ col_index_all = {}
30
+ for i_atlas in range(ModelType.n_atlas):
31
+ row_index = list(blocks_all[i_atlas]["spatial"][0])
32
+ col_index = list(blocks_all[i_atlas]["spatial"][1])
33
+ row_index_all[i_atlas] = torch.sort(torch.vstack(row_index).flatten())[
34
+ 0
35
+ ].tolist()
36
+ col_index_all[i_atlas] = torch.sort(torch.vstack(col_index).flatten())[
37
+ 0
38
+ ].tolist()
39
+
40
+ batch_features_all = [
41
+ torch.FloatTensor(feature_all[i][row_index_all[i], :].toarray()).to(device)
42
+ for i in range(ModelType.n_atlas)
43
+ ]
44
+
45
+ adj_all_block = [
46
+ torch.FloatTensor(
47
+ adj_all[i][row_index_all[i], :].tocsc()[:, col_index_all[i]].todense()
48
+ ).to(device)
49
+ if ModelType.input_identity[i] == "ST"
50
+ else model.scrna_seq_adj["atlas" + str(i)]()[row_index_all[i], :][
51
+ :, col_index_all[i]
52
+ ]
53
+ for i in range(ModelType.n_atlas)
54
+ ]
55
+ adj_all_block_dis = [
56
+ torch.FloatTensor(
57
+ adj_all[i][row_index_all[i], :].tocsc()[:, col_index_all[i]].todense()
58
+ ).to(device)
59
+ if ModelType.input_identity[i] == "ST"
60
+ else model.scrna_seq_adj["atlas" + str(i)]()[row_index_all[i], :][
61
+ :, col_index_all[i]
62
+ ].detach()
63
+ for i in range(ModelType.n_atlas)
64
+ ]
65
+
66
+ train_mask_batch_single = [
67
+ train_mask_i[row_index_all[blocks_all_ind]]
68
+ for train_mask_i, blocks_all_ind in zip(train_mask, blocks_all)
69
+ ]
70
+ train_mask_batch_spatial = [
71
+ train_mask_i[col_index_all[blocks_all_ind]]
72
+ for train_mask_i, blocks_all_ind in zip(train_mask, blocks_all)
73
+ ]
74
+
75
+ ### discriminator flags
76
+ flag_shape_single = [len(row_index_all[i]) for i in range(ModelType.n_atlas)]
77
+ flag_all_single = torch.cat(
78
+ [torch.full((x,), i) for i, x in enumerate(flag_shape_single)]
79
+ )
80
+ flag_source_cat_single = flag_all_single.long().to(device)
81
+
82
+ flag_shape_spatial = [len(col_index_all[i]) for i in range(ModelType.n_atlas)]
83
+ flag_all_spatial = torch.cat(
84
+ [torch.full((x,), i) for i, x in enumerate(flag_shape_spatial)]
85
+ )
86
+ flag_source_cat_spatial = flag_all_spatial.long().to(device)
87
+
88
+ return (
89
+ batch_features_all,
90
+ adj_all_block,
91
+ adj_all_block_dis,
92
+ train_mask_batch_single,
93
+ train_mask_batch_spatial,
94
+ flag_source_cat_single,
95
+ flag_source_cat_spatial,
96
+ row_index_all,
97
+ col_index_all,
98
+ )
99
+
100
+
101
+ def pretrain_model(
102
+ model,
103
+ spatial_dataloader,
104
+ feature_all,
105
+ adj_all,
106
+ device,
107
+ train_mask,
108
+ val_mask,
109
+ flagconfig,
110
+ ):
111
+ loss_atlas_val_best = float("inf")
112
+ patience_counter = 0
113
+
114
+ optimizer_dis = getattr(torch.optim, ModelType.optim_kw.value)(
115
+ itertools.chain(
116
+ model.discriminator_single.parameters(),
117
+ model.discriminator_spatial.parameters(),
118
+ ),
119
+ lr=ModelType.learning_rate.value,
120
+ )
121
+ optimizer_ae = getattr(torch.optim, ModelType.optim_kw.value)(
122
+ itertools.chain(
123
+ model.encoder.parameters(),
124
+ model.decoder.parameters(),
125
+ model.scrna_seq_adj.parameters(),
126
+ ),
127
+ lr=ModelType.learning_rate.value,
128
+ )
129
+ scheduler_dis = ReduceLROnPlateau(
130
+ optimizer_dis,
131
+ mode="min",
132
+ factor=ModelType.lr_factor_pretrain.value,
133
+ patience=ModelType.lr_patience_pretrain.value,
134
+ verbose=True,
135
+ )
136
+ scheduler_ae = ReduceLROnPlateau(
137
+ optimizer_ae,
138
+ mode="min",
139
+ factor=ModelType.lr_factor_pretrain.value,
140
+ patience=ModelType.lr_patience_pretrain.value,
141
+ verbose=True,
142
+ )
143
+
144
+ for epoch in tqdm(
145
+ range(ModelType.epochs_run_pretrain + 1, ModelType.n_epochs.value)
146
+ ):
147
+ loss_dis = 0
148
+ loss_ae_dis = 0
149
+ loss_all_item = 0
150
+ loss_atlas_i = {}
151
+ for i in range(ModelType.n_atlas):
152
+ loss_atlas_i[i] = 0
153
+ loss_atlas_val = 0
154
+ anneal = (
155
+ max(1 - (epoch - 1) / flagconfig.align_anneal, 0)
156
+ if flagconfig.align_anneal
157
+ else 0
158
+ )
159
+
160
+ model.train()
161
+
162
+ for blocks_all in spatial_dataloader:
163
+ (
164
+ batch_features_all,
165
+ adj_all_block,
166
+ adj_all_block_dis,
167
+ train_mask_batch_single,
168
+ train_mask_batch_spatial,
169
+ flag_source_cat_single,
170
+ flag_source_cat_spatial,
171
+ _,
172
+ _,
173
+ ) = get_data(blocks_all, feature_all, adj_all, train_mask, device, model)
174
+
175
+ # Train discriminator part
176
+ loss_part1 = compute_dis_loss_pretrain(
177
+ model,
178
+ flag_source_cat_single,
179
+ flag_source_cat_spatial,
180
+ anneal,
181
+ batch_features_all,
182
+ adj_all_block_dis,
183
+ train_mask_batch_single,
184
+ train_mask_batch_spatial,
185
+ flagconfig,
186
+ )
187
+ model.zero_grad(set_to_none=True)
188
+ loss_part1["dis"].backward()
189
+ optimizer_dis.step()
190
+ loss_dis += loss_part1["dis"].item()
191
+
192
+ # Train AE part
193
+ loss_part2 = compute_ae_loss_pretrain(
194
+ model,
195
+ flag_source_cat_single,
196
+ flag_source_cat_spatial,
197
+ anneal,
198
+ batch_features_all,
199
+ adj_all_block,
200
+ train_mask_batch_single,
201
+ train_mask_batch_spatial,
202
+ flagconfig,
203
+ )
204
+ model.zero_grad(set_to_none=True)
205
+ loss_part2["loss_all"].backward()
206
+ optimizer_ae.step()
207
+
208
+ if ModelType.use_llm_gene_embedding=='combine':
209
+ loss_part3 = compute_gene_embedding_loss(model)
210
+ model.zero_grad(set_to_none=True)
211
+ loss_part3.backward()
212
+ optimizer_ae.step()
213
+
214
+ for i in range(ModelType.n_atlas):
215
+ loss_atlas_i[i] += loss_part2["loss_AE_all"][i].item()
216
+ loss_all_item += loss_part2["loss_all"].item()
217
+ loss_ae_dis += loss_part2["dis_ae"].item()
218
+
219
+ flagconfig.align_anneal /= 2
220
+
221
+ if ModelType.verbose == True:
222
+ logging.info(
223
+ f"\n\nTrain Epoch {epoch}/{ModelType.n_epochs}, \
224
+ Loss dis: {loss_dis / len(spatial_dataloader)},\
225
+ Loss AE: {[i / len(spatial_dataloader) for i in loss_atlas_i.values()]} , \
226
+ Loss ae dis:{loss_ae_dis / len(spatial_dataloader)},\
227
+ Loss all:{loss_all_item / len(spatial_dataloader)}\n"
228
+ )
229
+
230
+ save_snapshot(model, epoch, ModelType.epochs_run_final, ModelType.snapshot_path,ModelType.verbose)
231
+
232
+ if not os.path.exists(f"{ModelType.save_dir}/lambda_disc_single.pkl"):
233
+ save_obj(
234
+ flagconfig.lambda_disc_single,
235
+ f"{ModelType.save_dir}/lambda_disc_single",
236
+ )
237
+
238
+ ################# validation
239
+ if epoch > ModelType.TRAIN_WITHOUT_EVAL.value:
240
+ model.eval()
241
+ with torch.no_grad():
242
+ for blocks_all in spatial_dataloader:
243
+
244
+ (
245
+ batch_features_all,
246
+ adj_all_block,
247
+ adj_all_block_dis,
248
+ val_mask_batch_single,
249
+ val_mask_batch_spatial,
250
+ flag_source_cat_single,
251
+ flag_source_cat_spatial,
252
+ _,
253
+ _,
254
+ ) = get_data(
255
+ blocks_all, feature_all, adj_all, val_mask, device, model
256
+ )
257
+
258
+ # val AE part
259
+ loss_part2 = compute_ae_loss_pretrain(
260
+ model,
261
+ flag_source_cat_single,
262
+ flag_source_cat_spatial,
263
+ anneal,
264
+ batch_features_all,
265
+ adj_all_block,
266
+ val_mask_batch_single,
267
+ val_mask_batch_spatial,
268
+ flagconfig,
269
+ )
270
+
271
+ for i in range(ModelType.n_atlas):
272
+ loss_atlas_val += loss_part2["loss_AE_all"][i].item()
273
+ # if np.isnan(loss_part2['loss_AE_all'][i].item()):
274
+ # p=0
275
+
276
+ loss_atlas_val = (
277
+ loss_atlas_val / len(spatial_dataloader) / ModelType.n_atlas
278
+ )
279
+
280
+ if ModelType.verbose == True:
281
+ logging.info(
282
+ f"\n\nValidation Epoch {epoch + 1}/{ModelType.n_epochs}, \
283
+ Loss AE validation: {loss_atlas_val} \n"
284
+ )
285
+
286
+ scheduler_dis.step(loss_atlas_val)
287
+ scheduler_ae.step(loss_atlas_val)
288
+ current_lr = optimizer_dis.param_groups[0]["lr"]
289
+ logging.info(f"\n\ncurrent lr:{current_lr}\n")
290
+
291
+ # If the loss is lower than the best loss so far, save the model And reset the patience counter
292
+ if loss_atlas_val < loss_atlas_val_best:
293
+ loss_atlas_val_best = loss_atlas_val
294
+ patience_counter = 0
295
+ torch.save(
296
+ model.state_dict(),
297
+ f"{ModelType.save_dir}/trained_model/FuseMap_pretrain_model.pt",
298
+ )
299
+
300
+ else:
301
+ patience_counter += 1
302
+
303
+ # If the patience counter is greater than or equal to the patience limit, stop training
304
+ if patience_counter >= ModelType.patience_limit_pretrain.value:
305
+ logging.info("\n\nEarly stopping due to loss not improving - patience count\n")
306
+ os.rename(
307
+ f"{ModelType.save_dir}/trained_model/FuseMap_pretrain_model.pt",
308
+ f"{ModelType.save_dir}/trained_model/FuseMap_pretrain_model_final.pt",
309
+ )
310
+ logging.info("\n\nFile name changed\n")
311
+ break
312
+ if current_lr < ModelType.lr_limit_pretrain.value:
313
+ logging.info("\n\nEarly stopping due to loss not improving - learning rate\n")
314
+ os.rename(
315
+ f"{ModelType.save_dir}/trained_model/FuseMap_pretrain_model.pt",
316
+ f"{ModelType.save_dir}/trained_model/FuseMap_pretrain_model_final.pt",
317
+ )
318
+ logging.info("\n\nFile name changed\n")
319
+ break
320
+
321
+ # torch.save(model.state_dict(), f"{ModelType.save_dir}/trained_model/FuseMap_pretrain_model_{epoch}.pt")
322
+
323
+ if os.path.exists(f"{ModelType.save_dir}/trained_model/FuseMap_pretrain_model.pt"):
324
+ os.rename(
325
+ f"{ModelType.save_dir}/trained_model/FuseMap_pretrain_model.pt",
326
+ f"{ModelType.save_dir}/trained_model/FuseMap_pretrain_model_final.pt",
327
+ )
328
+ logging.info("\n\nFile name changed in the end\n")
329
+
330
+
331
+ def train_model(
332
+ model,
333
+ spatial_dataloader,
334
+ feature_all,
335
+ adj_all,
336
+ device,
337
+ train_mask,
338
+ val_mask,
339
+ flagconfig,
340
+ ):
341
+ with open(f"{ModelType.save_dir}/balance_weight_single.pkl", "rb") as openfile:
342
+ balance_weight_single = pickle.load(openfile)
343
+ with open(f"{ModelType.save_dir}/balance_weight_spatial.pkl", "rb") as openfile:
344
+ balance_weight_spatial = pickle.load(openfile)
345
+ balance_weight_single = [i.to(device) for i in balance_weight_single]
346
+ balance_weight_spatial = [i.to(device) for i in balance_weight_spatial]
347
+
348
+ loss_atlas_val_best = float("inf")
349
+ patience_counter = 0
350
+
351
+ optimizer_dis = getattr(torch.optim, ModelType.optim_kw.value)(
352
+ itertools.chain(
353
+ model.discriminator_single.parameters(),
354
+ model.discriminator_spatial.parameters(),
355
+ ),
356
+ lr=ModelType.learning_rate.value,
357
+ )
358
+ optimizer_ae = getattr(torch.optim, ModelType.optim_kw.value)(
359
+ itertools.chain(
360
+ model.encoder.parameters(),
361
+ model.decoder.parameters(),
362
+ model.scrna_seq_adj.parameters(),
363
+ ),
364
+ lr=ModelType.learning_rate.value,
365
+ )
366
+ scheduler_dis = ReduceLROnPlateau(
367
+ optimizer_dis,
368
+ mode="min",
369
+ factor=ModelType.lr_factor_final.value,
370
+ patience=ModelType.lr_patience_final.value,
371
+ verbose=True,
372
+ )
373
+ scheduler_ae = ReduceLROnPlateau(
374
+ optimizer_ae,
375
+ mode="min",
376
+ factor=ModelType.lr_factor_final.value,
377
+ patience=ModelType.lr_patience_final.value,
378
+ verbose=True,
379
+ )
380
+
381
+ for epoch in tqdm(range(ModelType.epochs_run_final + 1, ModelType.n_epochs.value)):
382
+ loss_dis = 0
383
+ loss_ae_dis = 0
384
+ loss_all_item = 0
385
+ loss_atlas_i = {}
386
+ for i in range(ModelType.n_atlas):
387
+ loss_atlas_i[i] = 0
388
+ loss_atlas_val = 0
389
+ anneal = (
390
+ max(1 - (epoch - 1) / flagconfig.align_anneal, 0)
391
+ if flagconfig.align_anneal
392
+ else 0
393
+ )
394
+
395
+ model.train()
396
+
397
+ for blocks_all in spatial_dataloader:
398
+ (
399
+ batch_features_all,
400
+ adj_all_block,
401
+ adj_all_block_dis,
402
+ train_mask_batch_single,
403
+ train_mask_batch_spatial,
404
+ flag_source_cat_single,
405
+ flag_source_cat_spatial,
406
+ row_index_all,
407
+ col_index_all,
408
+ ) = get_data(blocks_all, feature_all, adj_all, train_mask, device, model)
409
+
410
+ balance_weight_single_block = [
411
+ balance_weight_single[i][row_index_all[i]]
412
+ for i in range(ModelType.n_atlas)
413
+ ]
414
+
415
+ balance_weight_spatial_block = [
416
+ balance_weight_spatial[i][col_index_all[i]]
417
+ for i in range(ModelType.n_atlas)
418
+ ]
419
+
420
+ # Train discriminator part
421
+ loss_part1 = compute_dis_loss(
422
+ model,
423
+ flag_source_cat_single,
424
+ flag_source_cat_spatial,
425
+ anneal,
426
+ batch_features_all,
427
+ adj_all_block_dis,
428
+ train_mask_batch_single,
429
+ train_mask_batch_spatial,
430
+ balance_weight_single_block,
431
+ balance_weight_spatial_block,
432
+ flagconfig,
433
+ )
434
+ model.zero_grad(set_to_none=True)
435
+ loss_part1["dis"].backward()
436
+ optimizer_dis.step()
437
+ loss_dis += loss_part1["dis"].item()
438
+
439
+ # Train AE part
440
+ loss_part2 = compute_ae_loss(
441
+ model,
442
+ flag_source_cat_single,
443
+ flag_source_cat_spatial,
444
+ anneal,
445
+ batch_features_all,
446
+ adj_all_block,
447
+ train_mask_batch_single,
448
+ train_mask_batch_spatial,
449
+ balance_weight_single_block,
450
+ balance_weight_spatial_block,
451
+ flagconfig,
452
+ )
453
+ model.zero_grad(set_to_none=True)
454
+ loss_part2["loss_all"].backward()
455
+ optimizer_ae.step()
456
+
457
+ if ModelType.use_llm_gene_embedding=='combine':
458
+ loss_part3 = compute_gene_embedding_loss(model)
459
+ model.zero_grad(set_to_none=True)
460
+ loss_part3.backward()
461
+ optimizer_ae.step()
462
+
463
+ for i in range(ModelType.n_atlas):
464
+ loss_atlas_i[i] += loss_part2["loss_AE_all"][i].item()
465
+ loss_all_item += loss_part2["loss_all"].item()
466
+ loss_ae_dis += loss_part2["dis_ae"].item()
467
+
468
+ flagconfig.align_anneal /= 2
469
+
470
+ save_snapshot(
471
+ model, ModelType.epochs_run_pretrain, epoch, ModelType.snapshot_path,ModelType.verbose
472
+ )
473
+
474
+ if ModelType.verbose == True:
475
+ logging.info(
476
+ f"\n\nTrain Epoch {epoch + 1}/{ModelType.n_epochs.value}, \
477
+ Loss dis: {loss_dis / len(spatial_dataloader)},\
478
+ Loss AE: {[i / len(spatial_dataloader) for i in loss_atlas_i.values()]} , \
479
+ Loss ae dis:{loss_ae_dis / len(spatial_dataloader)},\
480
+ Loss all:{loss_all_item / len(spatial_dataloader)}\n"
481
+ )
482
+
483
+ ################# validation
484
+ if epoch > ModelType.TRAIN_WITHOUT_EVAL.value:
485
+ model.eval()
486
+ with torch.no_grad():
487
+ for ind, blocks_all in enumerate(spatial_dataloader):
488
+ # if ind not in random_numbers:
489
+ # continue
490
+
491
+ (
492
+ batch_features_all,
493
+ adj_all_block,
494
+ adj_all_block_dis,
495
+ val_mask_batch_single,
496
+ val_mask_batch_spatial,
497
+ flag_source_cat_single,
498
+ flag_source_cat_spatial,
499
+ row_index_all,
500
+ col_index_all,
501
+ ) = get_data(
502
+ blocks_all, feature_all, adj_all, val_mask, device, model
503
+ )
504
+
505
+ balance_weight_single_block = [
506
+ balance_weight_single[i][row_index_all[i]]
507
+ for i in range(ModelType.n_atlas)
508
+ ]
509
+
510
+ balance_weight_spatial_block = [
511
+ balance_weight_spatial[i][col_index_all[i]]
512
+ for i in range(ModelType.n_atlas)
513
+ ]
514
+
515
+ # val AE part
516
+ loss_part2 = compute_ae_loss(
517
+ model,
518
+ flag_source_cat_single,
519
+ flag_source_cat_spatial,
520
+ anneal,
521
+ batch_features_all,
522
+ adj_all_block,
523
+ val_mask_batch_single,
524
+ val_mask_batch_spatial,
525
+ balance_weight_single_block,
526
+ balance_weight_spatial_block,
527
+ flagconfig,
528
+ )
529
+
530
+ for i in range(ModelType.n_atlas):
531
+ loss_atlas_val += loss_part2["loss_AE_all"][i].item()
532
+
533
+ loss_atlas_val = (
534
+ loss_atlas_val / len(spatial_dataloader) / ModelType.n_atlas
535
+ )
536
+ if ModelType.verbose == True:
537
+ logging.info(
538
+ f"\n\nValidation Epoch {epoch + 1}/{ModelType.n_epochs.value}, \
539
+ Loss AE validation: {loss_atlas_val} \n"
540
+ )
541
+
542
+ scheduler_dis.step(loss_atlas_val)
543
+ scheduler_ae.step(loss_atlas_val)
544
+ current_lr = optimizer_dis.param_groups[0]["lr"]
545
+ logging.info(f"\n\ncurrent lr:{current_lr}\n")
546
+
547
+ # If the loss is lower than the best loss so far, save the model And reset the patience counter
548
+ if loss_atlas_val < loss_atlas_val_best:
549
+ loss_atlas_val_best = loss_atlas_val
550
+ patience_counter = 0
551
+ torch.save(
552
+ model.state_dict(),
553
+ f"{ModelType.save_dir}/trained_model/FuseMap_final_model.pt",
554
+ )
555
+ else:
556
+ patience_counter += 1
557
+
558
+ # If the patience counter is greater than or equal to the patience limit, stop training
559
+ if patience_counter >= ModelType.patience_limit_final.value:
560
+ # torch.save(model.state_dict(), f"{save_dir}/trained_model/FuseMap_final_model_end.pt")
561
+ logging.info("\n\nEarly stopping due to loss not improving\n")
562
+ os.rename(
563
+ f"{ModelType.save_dir}/trained_model/FuseMap_final_model.pt",
564
+ f"{ModelType.save_dir}/trained_model/FuseMap_final_model_final.pt",
565
+ )
566
+ logging.info("\n\nFile name changed\n")
567
+ break
568
+ if current_lr < ModelType.lr_limit_final.value:
569
+ # torch.save(model.state_dict(), f"{save_dir}/trained_model/FuseMap_final_model_end.pt")
570
+ logging.info("\n\nEarly stopping due to loss not improving - learning rate\n")
571
+ os.rename(
572
+ f"{ModelType.save_dir}/trained_model/FuseMap_final_model.pt",
573
+ f"{ModelType.save_dir}/trained_model/FuseMap_final_model_final.pt",
574
+ )
575
+ logging.info("\n\nFile name changed\n")
576
+ break
577
+
578
+ if os.path.exists(f"{ModelType.save_dir}/trained_model/FuseMap_final_model.pt"):
579
+ os.rename(
580
+ f"{ModelType.save_dir}/trained_model/FuseMap_final_model.pt",
581
+ f"{ModelType.save_dir}/trained_model/FuseMap_final_model_final.pt",
582
+ )
583
+ logging.info("\n\nFile name changed in the end\n")
584
+
585
+
586
+ def read_model(
587
+ model, spatial_dataloader_test, g_all, feature_all, adj_all, device, ModelType, mode
588
+ ):
589
+ model.load_state_dict(
590
+ torch.load(f"{ModelType.save_dir}/trained_model/FuseMap_{mode}_model_final.pt")
591
+ )
592
+
593
+ with torch.no_grad():
594
+ model.eval()
595
+
596
+ learnt_scrna_seq_adj = {}
597
+ for i in range(ModelType.n_atlas):
598
+ if ModelType.input_identity[i] == "scrna":
599
+ learnt_scrna_seq_adj["atlas" + str(i)] = (
600
+ model.scrna_seq_adj["atlas" + str(i)]().detach().cpu().numpy()
601
+ )
602
+
603
+ for blocks_all in tqdm(spatial_dataloader_test):
604
+ row_index_all = {}
605
+ col_index_all = {}
606
+ for i_atlas in range(ModelType.n_atlas):
607
+ row_index = list(blocks_all[i_atlas]["spatial"][0])
608
+ col_index = list(blocks_all[i_atlas]["spatial"][1])
609
+ row_index_all[i_atlas] = torch.sort(torch.vstack(row_index).flatten())[
610
+ 0
611
+ ].tolist()
612
+ col_index_all[i_atlas] = torch.sort(torch.vstack(col_index).flatten())[
613
+ 0
614
+ ].tolist()
615
+
616
+ batch_features_all = [
617
+ torch.FloatTensor(feature_all[i][row_index_all[i], :].toarray()).to(
618
+ device
619
+ )
620
+ for i in range(ModelType.n_atlas)
621
+ ]
622
+ adj_all_block_dis = [
623
+ torch.FloatTensor(
624
+ adj_all[i][row_index_all[i], :]
625
+ .tocsc()[:, col_index_all[i]]
626
+ .todense()
627
+ ).to(device)
628
+ if ModelType.input_identity[i] == "ST"
629
+ else model.scrna_seq_adj["atlas" + str(i)]()[row_index_all[i], :][
630
+ :, col_index_all[i]
631
+ ].detach()
632
+ for i in range(ModelType.n_atlas)
633
+ ]
634
+
635
+ z_all = [
636
+ model.encoder["atlas" + str(i)](
637
+ batch_features_all[i], adj_all_block_dis[i]
638
+ )
639
+ for i in range(ModelType.n_atlas)
640
+ ]
641
+ z_distribution_all = [
642
+ z_all[i][3] for i in range(ModelType.n_atlas)
643
+ ]
644
+
645
+ z_spatial_all = [z_all[i][2] for i in range(ModelType.n_atlas)]
646
+
647
+ for i in range(ModelType.n_atlas):
648
+ g_all[i].nodes[row_index_all[i]].data["single_feat_hidden"] = (
649
+ z_distribution_all[i].detach().cpu()
650
+ )
651
+ g_all[i].nodes[col_index_all[i]].data["spatial_feat_hidden"] = (
652
+ z_spatial_all[i].detach().cpu()
653
+ )
654
+
655
+ latent_embeddings_all_single = [
656
+ g_all[i].ndata["single_feat_hidden"].numpy() for i in range(ModelType.n_atlas)
657
+ ]
658
+ latent_embeddings_all_spatial = [
659
+ g_all[i].ndata["spatial_feat_hidden"].numpy() for i in range(ModelType.n_atlas)
660
+ ]
661
+
662
+ save_obj(
663
+ latent_embeddings_all_single,
664
+ f"{ModelType.save_dir}/latent_embeddings_all_single_{mode}",
665
+ )
666
+ save_obj(
667
+ latent_embeddings_all_spatial,
668
+ f"{ModelType.save_dir}/latent_embeddings_all_spatial_{mode}",
669
+ )
670
+
671
+
672
+ def balance_weight(model, adatas, save_dir, n_atlas, device):
673
+ with open(
674
+ f"{save_dir}/latent_embeddings_all_single_pretrain.pkl", "rb"
675
+ ) as openfile:
676
+ latent_embeddings_all_single = pickle.load(openfile)
677
+ with open(
678
+ f"{save_dir}/latent_embeddings_all_spatial_pretrain.pkl", "rb"
679
+ ) as openfile:
680
+ latent_embeddings_all_spatial = pickle.load(openfile)
681
+
682
+ adatas_ = [
683
+ ad.AnnData(
684
+ obs=adatas[i].obs.copy(deep=False).assign(n=1),
685
+ obsm={
686
+ "single": latent_embeddings_all_single[i],
687
+ "spatial": latent_embeddings_all_spatial[i],
688
+ },
689
+ )
690
+ for i in range(n_atlas)
691
+ ]
692
+
693
+ if not os.path.exists(f"{save_dir}/ad_fusemap_single_leiden.pkl"):
694
+ leiden_adata_single = []
695
+ leiden_adata_spatial = []
696
+ ad_fusemap_single_leiden = []
697
+ ad_fusemap_spatial_leiden = []
698
+ for adata_ in adatas_:
699
+ sc.pp.neighbors(
700
+ adata_,
701
+ n_pcs=adata_.obsm["single"].shape[1],
702
+ use_rep="single",
703
+ metric="cosine",
704
+ )
705
+ sc.tl.leiden(adata_, resolution=1, key_added="fusemap_single_leiden")
706
+ ad_fusemap_single_leiden.append(list(adata_.obs["fusemap_single_leiden"]))
707
+ leiden_adata_single.append(
708
+ average_embeddings(adata_, "fusemap_single_leiden", "single")
709
+ )
710
+
711
+ sc.pp.neighbors(
712
+ adata_,
713
+ n_pcs=adata_.obsm["spatial"].shape[1],
714
+ use_rep="spatial",
715
+ metric="cosine",
716
+ )
717
+ sc.tl.leiden(adata_, resolution=1, key_added="fusemap_spatial_leiden")
718
+ ad_fusemap_spatial_leiden.append(list(adata_.obs["fusemap_spatial_leiden"]))
719
+ leiden_adata_spatial.append(
720
+ average_embeddings(adata_, "fusemap_spatial_leiden", "spatial")
721
+ )
722
+
723
+ save_obj(ad_fusemap_single_leiden, f"{save_dir}/ad_fusemap_single_leiden")
724
+ save_obj(ad_fusemap_spatial_leiden, f"{save_dir}/ad_fusemap_spatial_leiden")
725
+ save_obj(leiden_adata_single, f"{save_dir}/leiden_adata_single")
726
+ save_obj(leiden_adata_spatial, f"{save_dir}/leiden_adata_spatial")
727
+
728
+ else:
729
+ with open(f"{save_dir}/ad_fusemap_single_leiden.pkl", "rb") as openfile:
730
+ ad_fusemap_single_leiden = pickle.load(openfile)
731
+ with open(f"{save_dir}/ad_fusemap_spatial_leiden.pkl", "rb") as openfile:
732
+ ad_fusemap_spatial_leiden = pickle.load(openfile)
733
+ try:
734
+ with open(f"{save_dir}/leiden_adata_single.pkl", "rb") as openfile:
735
+ leiden_adata_single = pickle.load(openfile)
736
+ with open(f"{save_dir}/leiden_adata_spatial.pkl", "rb") as openfile:
737
+ leiden_adata_spatial = pickle.load(openfile)
738
+ except:
739
+ ### need to convert
740
+ leiden_adata_single = []
741
+ for i in range(len(ad_fusemap_single_leiden)):
742
+ leiden_adata_single.append(
743
+ sc.read_h5ad(
744
+ f"{save_dir}/pickle_convert/PRETRAINED_leiden_adata_single_{i}.h5ad"
745
+ )
746
+ )
747
+ leiden_adata_spatial = []
748
+ for i in range(len(ad_fusemap_single_leiden)):
749
+ leiden_adata_spatial.append(
750
+ sc.read_h5ad(
751
+ f"{save_dir}/pickle_convert/PRETRAINED_leiden_adata_spatial_{i}.h5ad"
752
+ )
753
+ )
754
+
755
+ for ind, adata_ in enumerate(adatas_):
756
+ adata_.obs["fusemap_single_leiden"] = ad_fusemap_single_leiden[ind]
757
+ adata_.obs["fusemap_spatial_leiden"] = ad_fusemap_spatial_leiden[ind]
758
+
759
+ if len(leiden_adata_single) > 10:
760
+ # raise ValueError('balance weight')
761
+ balance_weight_single = get_balance_weight_subsample(
762
+ leiden_adata_single, adatas_, "fusemap_single_leiden"
763
+ )
764
+ balance_weight_spatial = get_balance_weight_subsample(
765
+ leiden_adata_spatial, adatas_, "fusemap_spatial_leiden"
766
+ )
767
+ else:
768
+ balance_weight_single = get_balance_weight(
769
+ adatas, leiden_adata_single, adatas_, "fusemap_single_leiden"
770
+ )
771
+ balance_weight_spatial = get_balance_weight(
772
+ adatas, leiden_adata_spatial, adatas_, "fusemap_spatial_leiden"
773
+ )
774
+
775
+ balance_weight_single = [torch.tensor(i).to(device) for i in balance_weight_single]
776
+ balance_weight_spatial = [
777
+ torch.tensor(i).to(device) for i in balance_weight_spatial
778
+ ]
779
+
780
+ save_obj(balance_weight_single, f"{save_dir}/balance_weight_single")
781
+ save_obj(balance_weight_spatial, f"{save_dir}/balance_weight_spatial")
782
+
783
+
784
+ def load_ref_model(
785
+ ref_dir,
786
+ device,
787
+ ):
788
+ PRETRAINED_MODEL_PATH = ref_dir + f"/pretrain_model.pt"
789
+
790
+ if os.path.exists(PRETRAINED_MODEL_PATH):
791
+ TRAINED_MODEL = torch.load(PRETRAINED_MODEL_PATH, map_location=device)
792
+ TRAINED_X_NUM = sum(["decoder" in i for i in TRAINED_MODEL.keys()])
793
+
794
+ TRAINED_GENE_EMBED = sc.read_h5ad(ref_dir + "/ad_gene_embedding.h5ad")
795
+ TRAINED_GENE_NAME = list(TRAINED_GENE_EMBED.obs.index)
796
+ else:
797
+ raise ValueError("No pretrained model found!")
798
+ return TRAINED_MODEL, TRAINED_X_NUM, TRAINED_GENE_EMBED, TRAINED_GENE_NAME
799
+
800
+
801
+ def add_pretrain_to_name(s):
802
+ if "discriminator_single" in s:
803
+ return s.replace("discriminator_single", "discriminator_single_pretrain")
804
+ elif "discriminator_spatial" in s:
805
+ return s.replace("discriminator_spatial", "discriminator_spatial_pretrain")
806
+ else:
807
+ return s
808
+
809
+
810
+ def transfer_weight(TRAINED_MODEL, pretrain_index, adapt_model):
811
+ layers_to_transfer = [
812
+ "discriminator_single.linear_0.weight",
813
+ "discriminator_single.linear_0.bias",
814
+ "discriminator_single.linear_1.weight",
815
+ "discriminator_single.linear_1.bias",
816
+ "discriminator_spatial.linear_0.weight",
817
+ "discriminator_spatial.linear_0.bias",
818
+ "discriminator_spatial.linear_1.weight",
819
+ "discriminator_spatial.linear_1.bias",
820
+ ]
821
+ transferred_dict = {
822
+ k: v for k, v in TRAINED_MODEL.items() if k in layers_to_transfer
823
+ }
824
+ transferred_dict_pretrain = {
825
+ add_pretrain_to_name(k): v
826
+ for k, v in TRAINED_MODEL.items()
827
+ if k in layers_to_transfer
828
+ }
829
+ transferred_dict.update(transferred_dict_pretrain)
830
+
831
+ new_model_dict = adapt_model.state_dict()
832
+ new_model_dict.update(transferred_dict)
833
+ adapt_model.load_state_dict(new_model_dict)
834
+
835
+ with torch.no_grad():
836
+ # Assuming the pretrained parameters go into the first 'n' units
837
+ adapt_model.discriminator_single_pretrain.pred.weight = nn.Parameter(
838
+ TRAINED_MODEL["discriminator_single.pred.weight"]
839
+ )
840
+ adapt_model.discriminator_single_pretrain.pred.bias = nn.Parameter(
841
+ TRAINED_MODEL["discriminator_single.pred.bias"]
842
+ )
843
+ adapt_model.discriminator_spatial_pretrain.pred.weight = nn.Parameter(
844
+ TRAINED_MODEL["discriminator_spatial.pred.weight"]
845
+ )
846
+ adapt_model.discriminator_spatial_pretrain.pred.bias = nn.Parameter(
847
+ TRAINED_MODEL["discriminator_spatial.pred.bias"]
848
+ )
849
+ adapt_model.gene_embedding_pretrained = nn.Parameter(
850
+ TRAINED_MODEL["gene_embedding"][:, pretrain_index]
851
+ )
852
+
853
+ for param in adapt_model.discriminator_single_pretrain.parameters():
854
+ param.requires_grad = False
855
+ for param in adapt_model.discriminator_spatial_pretrain.parameters():
856
+ param.requires_grad = False
857
+ adapt_model.gene_embedding_pretrained.requires_grad = False
858
+
859
+ adapt_model.discriminator_single.pred.weight.requires_grad = True
860
+ adapt_model.discriminator_single.pred.bias.requires_grad = True
861
+ adapt_model.discriminator_spatial.pred.weight.requires_grad = True
862
+ adapt_model.discriminator_spatial.pred.bias.requires_grad = True
863
+
864
+ # Print out to verify
865
+ # for name, param in adapt_model.named_parameters():
866
+ # print(name, param.requires_grad)
867
+
868
+
869
+ def load_ref_data(ref_dir, TRAINED_X_NUM, batch_size, USE_REFERENCE_PCT=0.1):
870
+ with open(ref_dir + f"/latent_embeddings_single.pkl", "rb") as openfile:
871
+ latent_embeddings_single = pickle.load(openfile)
872
+ with open(ref_dir + f"/latent_embeddings_spatial.pkl", "rb") as openfile:
873
+ latent_embeddings_spatial = pickle.load(openfile)
874
+
875
+ ds_pretrain_single = [
876
+ MapPretrainDataset(latent_embeddings_single[i]) for i in range(TRAINED_X_NUM)
877
+ ]
878
+ dataloader_pretrain_single = MapPretrainDataLoader(
879
+ ds_pretrain_single,
880
+ int(batch_size * USE_REFERENCE_PCT * 4),
881
+ shuffle=True,
882
+ n_atlas=TRAINED_X_NUM,
883
+ )
884
+
885
+ ds_pretrain_spatial = [
886
+ MapPretrainDataset(latent_embeddings_spatial[i]) for i in range(TRAINED_X_NUM)
887
+ ]
888
+ dataloader_pretrain_spatial = MapPretrainDataLoader(
889
+ ds_pretrain_spatial,
890
+ int(batch_size * USE_REFERENCE_PCT),
891
+ shuffle=True,
892
+ n_atlas=TRAINED_X_NUM,
893
+ )
894
+
895
+ return dataloader_pretrain_single, dataloader_pretrain_spatial
896
+
897
+
898
+ def map_model(
899
+ adapt_model,
900
+ spatial_dataloader,
901
+ feature_all,
902
+ adj_all,
903
+ device,
904
+ train_mask,
905
+ val_mask,
906
+ ref_dir,
907
+ dataloader_pretrain_single,
908
+ dataloader_pretrain_spatial,
909
+ TRAINED_X_NUM,
910
+ flagconfig,
911
+ ):
912
+ loss_atlas_val_best = float("inf")
913
+ patience_counter = 0
914
+
915
+ optimizer_dis = getattr(torch.optim, ModelType.optim_kw.value)(
916
+ itertools.chain(
917
+ adapt_model.discriminator_single.parameters(),
918
+ adapt_model.discriminator_spatial.parameters(),
919
+ ),
920
+ lr=ModelType.learning_rate.value,
921
+ )
922
+ optimizer_ae = getattr(torch.optim, ModelType.optim_kw.value)(
923
+ itertools.chain(
924
+ adapt_model.encoder.parameters(),
925
+ adapt_model.decoder.parameters(),
926
+ adapt_model.scrna_seq_adj.parameters(),
927
+ ),
928
+ lr=ModelType.learning_rate.value,
929
+ )
930
+ scheduler_dis = ReduceLROnPlateau(
931
+ optimizer_dis,
932
+ mode="min",
933
+ factor=ModelType.lr_factor_pretrain.value,
934
+ patience=ModelType.lr_patience_pretrain.value,
935
+ verbose=True,
936
+ )
937
+ scheduler_ae = ReduceLROnPlateau(
938
+ optimizer_ae,
939
+ mode="min",
940
+ factor=ModelType.lr_factor_pretrain.value,
941
+ patience=ModelType.lr_patience_pretrain.value,
942
+ verbose=True,
943
+ )
944
+
945
+ dataloader_pretrain_single_cycle = itertools.cycle(dataloader_pretrain_single)
946
+ dataloader_pretrain_spatial_cycle = itertools.cycle(dataloader_pretrain_spatial)
947
+
948
+ for epoch in tqdm(
949
+ range(ModelType.epochs_run_pretrain + 1, ModelType.n_epochs.value)
950
+ ):
951
+ loss_dis = 0
952
+ loss_ae_dis = 0
953
+ loss_all_item = 0
954
+ loss_atlas_i = {}
955
+ for i in range(ModelType.n_atlas):
956
+ loss_atlas_i[i] = 0
957
+ loss_atlas_val = 0
958
+ anneal = (
959
+ max(1 - (epoch - 1) / flagconfig.align_anneal, 0)
960
+ if flagconfig.align_anneal
961
+ else 0
962
+ )
963
+
964
+ adapt_model.train()
965
+
966
+ for blocks_all in tqdm(spatial_dataloader):
967
+ (
968
+ batch_features_all,
969
+ adj_all_block,
970
+ adj_all_block_dis,
971
+ train_mask_batch_single,
972
+ train_mask_batch_spatial,
973
+ flag_source_cat_single,
974
+ flag_source_cat_spatial,
975
+ _,
976
+ _,
977
+ ) = get_data(
978
+ blocks_all, feature_all, adj_all, train_mask, device, adapt_model
979
+ )
980
+
981
+ ### difference: add pretrain data
982
+ pretrain_single_batch = next(dataloader_pretrain_single_cycle)
983
+ pretrain_single_batch = [
984
+ pretrain_single_batch[i].to(device) for i in range(TRAINED_X_NUM)
985
+ ]
986
+ pretrain_spatial_batch = next(dataloader_pretrain_spatial_cycle)
987
+ pretrain_spatial_batch = [
988
+ pretrain_spatial_batch[i].to(device) for i in range(TRAINED_X_NUM)
989
+ ]
990
+
991
+ ### add difference: add discriminator pretrain
992
+ flag_shape_single_pretrain = [
993
+ pretrain_single_batch[i].shape[0] for i in range(TRAINED_X_NUM)
994
+ ]
995
+ flag_all_single_pretrain = torch.cat(
996
+ [
997
+ torch.full((x,), i + ModelType.n_atlas)
998
+ for i, x in enumerate(flag_shape_single_pretrain)
999
+ ]
1000
+ )
1001
+ flag_source_cat_single_pretrain = flag_all_single_pretrain.long().to(device)
1002
+
1003
+ flag_shape_spatial_pretrain = [
1004
+ pretrain_spatial_batch[i].shape[0] for i in range(TRAINED_X_NUM)
1005
+ ]
1006
+ flag_all_spatial_pretrain = torch.cat(
1007
+ [
1008
+ torch.full((x,), i + ModelType.n_atlas)
1009
+ for i, x in enumerate(flag_shape_spatial_pretrain)
1010
+ ]
1011
+ )
1012
+ flag_source_cat_spatial_pretrain = flag_all_spatial_pretrain.long().to(
1013
+ device
1014
+ )
1015
+
1016
+ # Train discriminator part
1017
+ loss_part1 = compute_dis_loss_map(
1018
+ adapt_model,
1019
+ flag_source_cat_single,
1020
+ flag_source_cat_spatial,
1021
+ anneal,
1022
+ batch_features_all,
1023
+ adj_all_block_dis,
1024
+ train_mask_batch_single,
1025
+ train_mask_batch_spatial,
1026
+ pretrain_single_batch,
1027
+ pretrain_spatial_batch,
1028
+ flag_source_cat_single_pretrain,
1029
+ flag_source_cat_spatial_pretrain,
1030
+ flagconfig,
1031
+ )
1032
+ adapt_model.zero_grad(set_to_none=True)
1033
+ loss_part1["dis"].backward()
1034
+ optimizer_dis.step()
1035
+ loss_dis += loss_part1["dis"].item()
1036
+
1037
+ # Train AE part
1038
+ loss_part2 = compute_ae_loss_map(
1039
+ adapt_model,
1040
+ flag_source_cat_single,
1041
+ flag_source_cat_spatial,
1042
+ anneal,
1043
+ batch_features_all,
1044
+ adj_all_block,
1045
+ train_mask_batch_single,
1046
+ train_mask_batch_spatial,
1047
+ pretrain_single_batch,
1048
+ pretrain_spatial_batch,
1049
+ flag_source_cat_single_pretrain,
1050
+ flag_source_cat_spatial_pretrain,
1051
+ flagconfig,
1052
+ )
1053
+ adapt_model.zero_grad(set_to_none=True)
1054
+ loss_part2["loss_all"].backward()
1055
+ optimizer_ae.step()
1056
+
1057
+ for i in range(ModelType.n_atlas):
1058
+ loss_atlas_i[i] += loss_part2["loss_AE_all"][i].item()
1059
+ loss_all_item += loss_part2["loss_all"].item()
1060
+ loss_ae_dis += loss_part2["dis_ae"].item()
1061
+
1062
+ flagconfig.align_anneal /= 2
1063
+
1064
+ if ModelType.verbose == True:
1065
+ logging.info(
1066
+ f"\n\nTrain Epoch {epoch}/{ModelType.n_epochs}, \
1067
+ Loss dis: {loss_dis / len(spatial_dataloader)},\
1068
+ Loss AE: {[i / len(spatial_dataloader) for i in loss_atlas_i.values()]} , \
1069
+ Loss ae dis:{loss_ae_dis / len(spatial_dataloader)},\
1070
+ Loss all:{loss_all_item / len(spatial_dataloader)}\n"
1071
+ )
1072
+
1073
+ save_snapshot(
1074
+ adapt_model, epoch, ModelType.epochs_run_final, ModelType.snapshot_path, ModelType.verbose
1075
+ )
1076
+
1077
+ if not os.path.exists(f"{ModelType.save_dir}/lambda_disc_single.pkl"):
1078
+ save_obj(
1079
+ flagconfig.lambda_disc_single,
1080
+ f"{ModelType.save_dir}/lambda_disc_single",
1081
+ )
1082
+
1083
+ ################# validation
1084
+ if epoch > ModelType.TRAIN_WITHOUT_EVAL.value:
1085
+ adapt_model.eval()
1086
+ with torch.no_grad():
1087
+ for blocks_all in spatial_dataloader:
1088
+ (
1089
+ batch_features_all,
1090
+ adj_all_block,
1091
+ adj_all_block_dis,
1092
+ val_mask_batch_single,
1093
+ val_mask_batch_spatial,
1094
+ flag_source_cat_single,
1095
+ flag_source_cat_spatial,
1096
+ _,
1097
+ _,
1098
+ ) = get_data(
1099
+ blocks_all, feature_all, adj_all, val_mask, device, adapt_model
1100
+ )
1101
+
1102
+ ### difference: add pretrain data
1103
+ pretrain_single_batch = next(dataloader_pretrain_single_cycle)
1104
+ pretrain_single_batch = [
1105
+ pretrain_single_batch[i].to(device)
1106
+ for i in range(TRAINED_X_NUM)
1107
+ ]
1108
+ pretrain_spatial_batch = next(dataloader_pretrain_spatial_cycle)
1109
+ pretrain_spatial_batch = [
1110
+ pretrain_spatial_batch[i].to(device)
1111
+ for i in range(TRAINED_X_NUM)
1112
+ ]
1113
+
1114
+ ### difference: discriminator pretrain
1115
+ flag_shape_single_pretrain = [
1116
+ pretrain_single_batch[i].shape[0] for i in range(TRAINED_X_NUM)
1117
+ ]
1118
+ flag_all_single_pretrain = torch.cat(
1119
+ [
1120
+ torch.full((x,), i + ModelType.n_atlas)
1121
+ for i, x in enumerate(flag_shape_single_pretrain)
1122
+ ]
1123
+ )
1124
+ flag_source_cat_single_pretrain = (
1125
+ flag_all_single_pretrain.long().to(device)
1126
+ )
1127
+
1128
+ flag_shape_spatial_pretrain = [
1129
+ pretrain_spatial_batch[i].shape[0] for i in range(TRAINED_X_NUM)
1130
+ ]
1131
+ flag_all_spatial_pretrain = torch.cat(
1132
+ [
1133
+ torch.full((x,), i + ModelType.n_atlas)
1134
+ for i, x in enumerate(flag_shape_spatial_pretrain)
1135
+ ]
1136
+ )
1137
+ flag_source_cat_spatial_pretrain = (
1138
+ flag_all_spatial_pretrain.long().to(device)
1139
+ )
1140
+
1141
+ # val AE part
1142
+ loss_part2 = compute_ae_loss_map(
1143
+ adapt_model,
1144
+ flag_source_cat_single,
1145
+ flag_source_cat_spatial,
1146
+ anneal,
1147
+ batch_features_all,
1148
+ adj_all_block,
1149
+ val_mask_batch_single,
1150
+ val_mask_batch_spatial,
1151
+ pretrain_single_batch,
1152
+ pretrain_spatial_batch,
1153
+ flag_source_cat_single_pretrain,
1154
+ flag_source_cat_spatial_pretrain,
1155
+ flagconfig,
1156
+ )
1157
+
1158
+ for i in range(ModelType.n_atlas):
1159
+ loss_atlas_val += loss_part2["loss_AE_all"][i].item()
1160
+
1161
+ loss_atlas_val = loss_atlas_val / len(spatial_dataloader) / ModelType.n_atlas
1162
+ if ModelType.verbose == True:
1163
+ logging.info(
1164
+ f"\n\nValidation Epoch {epoch + 1}/{ModelType.n_epochs.value}, \
1165
+ Loss AE validation: {loss_atlas_val} \n"
1166
+ )
1167
+
1168
+ scheduler_dis.step(loss_atlas_val)
1169
+ scheduler_ae.step(loss_atlas_val)
1170
+ current_lr = optimizer_dis.param_groups[0]["lr"]
1171
+ logging.info(f"\n\ncurrent lr:{current_lr}\n")
1172
+
1173
+ # If the loss is lower than the best loss so far, save the model And reset the patience counter
1174
+ if loss_atlas_val < loss_atlas_val_best:
1175
+ loss_atlas_val_best = loss_atlas_val
1176
+ patience_counter = 0
1177
+ torch.save(
1178
+ adapt_model.state_dict(),
1179
+ f"{ModelType.save_dir}/trained_model/FuseMap_map_model.pt",
1180
+ )
1181
+
1182
+ else:
1183
+ patience_counter += 1
1184
+
1185
+
1186
+ # If the patience counter is greater than or equal to the patience limit, stop training
1187
+ if patience_counter >= ModelType.patience_limit_final.value:
1188
+ # torch.save(model.state_dict(), f"{save_dir}/trained_model/FuseMap_final_model_end.pt")
1189
+ logging.info("\n\nEarly stopping due to loss not improving\n")
1190
+ os.rename(
1191
+ f"{ModelType.save_dir}/trained_model/FuseMap_map_model.pt",
1192
+ f"{ModelType.save_dir}/trained_model/FuseMap_map_model_final.pt",
1193
+ )
1194
+ logging.info("\n\nFile name changed\n")
1195
+ break
1196
+ if current_lr < ModelType.lr_limit_final.value:
1197
+ # torch.save(model.state_dict(), f"{save_dir}/trained_model/FuseMap_final_model_end.pt")
1198
+ logging.info("\n\nEarly stopping due to loss not improving - learning rate\n")
1199
+ os.rename(
1200
+ f"{ModelType.save_dir}/trained_model/FuseMap_map_model.pt",
1201
+ f"{ModelType.save_dir}/trained_model/FuseMap_map_model_final.pt",
1202
+ )
1203
+ logging.info("\n\nFile name changed\n")
1204
+ break
1205
+
1206
+ if os.path.exists(f"{ModelType.save_dir}/trained_model/FuseMap_map_model.pt"):
1207
+ os.rename(
1208
+ f"{ModelType.save_dir}/trained_model/FuseMap_map_model.pt",
1209
+ f"{ModelType.save_dir}/trained_model/FuseMap_map_model_final.pt",
1210
+ )
1211
+ logging.info("\n\nFile name changed in the end\n")