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/loss.py ADDED
@@ -0,0 +1,843 @@
1
+ import logging
2
+ import torch.nn.functional as F
3
+ import sklearn
4
+ import numpy as np
5
+ import torch
6
+ import torch.distributions as D
7
+ import pandas as pd
8
+ from sparse import COO
9
+ from fusemap.config import *
10
+ import torch.nn as nn
11
+
12
+ def AE_Gene_loss(recon_x, x, z_distribution):
13
+ if recon_x.shape[0] == 0:
14
+ return torch.tensor(0.0, dtype=torch.float32).to(recon_x.device)
15
+
16
+ reconstruction_loss = F.mse_loss(recon_x, x)
17
+ kl_divergence = (
18
+ D.kl_divergence(z_distribution, D.Normal(0.0, 1.0)).sum(dim=1).mean()
19
+ / x.shape[1]
20
+ )
21
+ return reconstruction_loss + kl_divergence
22
+
23
+
24
+ def prod(x):
25
+ ########### function from GLUE: https://github.com/gao-lab/GLUE
26
+ # try:
27
+ # from math import prod # pylint: disable=redefined-outer-name
28
+ # return np.prod(x)
29
+ # except ImportError:
30
+ ans = 1
31
+ for item in x:
32
+ ans = ans * item
33
+ return ans
34
+
35
+
36
+ """
37
+ pretrain loss
38
+ """
39
+
40
+
41
+ def compute_gene_embedding_loss(
42
+ model
43
+ ):
44
+ # Calculate gene embedding loss
45
+ learned_matrix = model.gene_embedding.T
46
+
47
+ learned_matrix=learned_matrix[model.llm_ind,:]
48
+
49
+ learned_matrix_normalized = learned_matrix / learned_matrix.norm(dim=1, keepdim=True)
50
+ predicted_matrix = torch.matmul(learned_matrix_normalized, learned_matrix_normalized.T)
51
+
52
+ loss_fn = nn.MSELoss()
53
+ loss_part3 = loss_fn(predicted_matrix, model.ground_truth_rel_matrix)
54
+ return loss_part3
55
+
56
+
57
+
58
+ def compute_dis_loss_pretrain(
59
+ model,
60
+ flag_source_cat_single,
61
+ flag_source_cat_spatial,
62
+ anneal,
63
+ batch_features_all,
64
+ adj_all,
65
+ mask_batch_single,
66
+ mask_batch_spatial,
67
+ flagconfig,
68
+ ):
69
+ mask_batch_single_all = torch.hstack(mask_batch_single)
70
+ mask_batch_spatial_all = torch.hstack(mask_batch_spatial)
71
+
72
+ z_all = [
73
+ model.encoder["atlas" + str(i)](batch_features_all[i], adj_all[i])
74
+ for i in range(ModelType.n_atlas)
75
+ ]
76
+ z_mean_cat_single = torch.cat([z_all[i][3] for i in range(ModelType.n_atlas)])[
77
+ mask_batch_single_all, :
78
+ ]
79
+
80
+ z_spatial_all = [z_all[i][2] for i in range(ModelType.n_atlas)]
81
+ z_mean_cat_spatial = torch.cat(z_spatial_all)[mask_batch_spatial_all, :]
82
+
83
+ if anneal:
84
+ if z_mean_cat_single.shape[0] > 1:
85
+ noise_single = D.Normal(0, z_mean_cat_single.std(axis=0)).sample(
86
+ (z_mean_cat_single.shape[0],)
87
+ )
88
+ z_mean_cat_single = (
89
+ z_mean_cat_single
90
+ + (anneal * ModelType.align_noise_coef.value) * noise_single
91
+ )
92
+ if z_mean_cat_spatial.shape[0] > 1:
93
+ noise_spatial = D.Normal(
94
+ 0, ModelType.EPS.value + z_mean_cat_spatial.std(axis=0)
95
+ ).sample((z_mean_cat_spatial.shape[0],))
96
+ z_mean_cat_spatial = (
97
+ z_mean_cat_spatial
98
+ + (anneal * ModelType.align_noise_coef.value) * noise_spatial
99
+ )
100
+
101
+ ### compute dis loss
102
+ loss_dis_single = F.cross_entropy(
103
+ F.softmax(model.discriminator_single(z_mean_cat_single), dim=1),
104
+ flag_source_cat_single[mask_batch_single_all],
105
+ reduction="none",
106
+ )
107
+ loss_dis_single = loss_dis_single.sum() / loss_dis_single.numel()
108
+
109
+ loss_dis_spatial = F.cross_entropy(
110
+ F.softmax(model.discriminator_spatial(z_mean_cat_spatial), dim=1),
111
+ flag_source_cat_spatial[mask_batch_spatial_all],
112
+ reduction="none",
113
+ )
114
+ loss_dis_spatial = loss_dis_spatial.sum() / loss_dis_spatial.numel()
115
+
116
+ loss_dis = flagconfig.lambda_disc_single * (loss_dis_single + loss_dis_spatial)
117
+
118
+ loss_all = {"dis": loss_dis}
119
+ return loss_all
120
+
121
+
122
+ def compute_ae_loss_pretrain(
123
+ model,
124
+ flag_source_cat_single,
125
+ flag_source_cat_spatial,
126
+ anneal,
127
+ batch_features_all,
128
+ adj_all,
129
+ mask_batch_single,
130
+ mask_batch_spatial,
131
+ flagconfig,
132
+ ):
133
+ z_all = [
134
+ model.encoder["atlas" + str(i)](batch_features_all[i], adj_all[i])
135
+ for i in range(ModelType.n_atlas)
136
+ ]
137
+
138
+ z_spatial_all = [z_all[i][2] for i in range(ModelType.n_atlas)]
139
+
140
+ # z_sample_all[i],
141
+
142
+ decoder_all = [
143
+ model.decoder["atlas" + str(i)](z_spatial_all[i], adj_all[i])
144
+ for i in range(ModelType.n_atlas)
145
+ ]
146
+
147
+ z_distribution_loss = [
148
+ z_all[i][0]
149
+
150
+ for i in range(ModelType.n_atlas)
151
+ ]
152
+ loss_AE_all = [
153
+ ModelType.lambda_ae_single.value
154
+ * AE_Gene_loss(
155
+ decoder_all[i][mask_batch_single[i], :],
156
+ batch_features_all[i][mask_batch_single[i], :],
157
+ z_distribution_loss[i],
158
+ )
159
+ for i in range(ModelType.n_atlas)
160
+ ]
161
+
162
+ mask_batch_single_all = torch.hstack(mask_batch_single)
163
+ mask_batch_spatial_all = torch.hstack(mask_batch_spatial)
164
+
165
+ z_mean_cat_single = torch.cat([z_all[i][3] for i in range(ModelType.n_atlas)])[
166
+ mask_batch_single_all, :
167
+ ]
168
+ z_mean_cat_spatial = torch.cat(z_spatial_all)[mask_batch_spatial_all, :]
169
+
170
+ if anneal:
171
+ if z_mean_cat_single.shape[0] > 1:
172
+ noise_single = D.Normal(0, z_mean_cat_single.std(axis=0)).sample(
173
+ (z_mean_cat_single.shape[0],)
174
+ )
175
+ z_mean_cat_single = (
176
+ z_mean_cat_single
177
+ + (anneal * ModelType.align_noise_coef.value) * noise_single
178
+ )
179
+ if z_mean_cat_spatial.shape[0] > 1:
180
+ noise_spatial = D.Normal(
181
+ 0, ModelType.EPS.value + z_mean_cat_spatial.std(axis=0)
182
+ ).sample((z_mean_cat_spatial.shape[0],))
183
+ z_mean_cat_spatial = (
184
+ z_mean_cat_spatial
185
+ + (anneal * ModelType.align_noise_coef.value) * noise_spatial
186
+ )
187
+
188
+ ### compute dis loss
189
+
190
+ loss_dis_single = F.cross_entropy(
191
+ F.softmax(model.discriminator_single(z_mean_cat_single), dim=1),
192
+ flag_source_cat_single[mask_batch_single_all],
193
+ reduction="none",
194
+ )
195
+ loss_dis_single = loss_dis_single.sum() / loss_dis_single.numel()
196
+
197
+ loss_dis_spatial = F.cross_entropy(
198
+ F.softmax(model.discriminator_spatial(z_mean_cat_spatial), dim=1),
199
+ flag_source_cat_spatial[mask_batch_spatial_all],
200
+ reduction="none",
201
+ )
202
+ loss_dis_spatial = loss_dis_spatial.sum() / loss_dis_spatial.numel()
203
+
204
+ loss_dis = flagconfig.lambda_disc_single * (loss_dis_single + loss_dis_spatial)
205
+
206
+ if (
207
+ flagconfig.lambda_disc_single == 1
208
+ ): # and loss_dis.item()<sum(loss_AE_all).item()/DIS_LAMDA:
209
+ flagconfig.lambda_disc_single = (
210
+ sum(loss_AE_all).item() / ModelType.DIS_LAMDA.value / loss_dis.item()
211
+ )
212
+ print(f"lambda_disc_single changed to {flagconfig.lambda_disc_single}")
213
+ loss_dis = flagconfig.lambda_disc_single * loss_dis
214
+
215
+ loss_all = {
216
+ "dis_ae": loss_dis,
217
+ "loss_AE_all": loss_AE_all,
218
+ "loss_all": -loss_dis + sum(loss_AE_all),
219
+ }
220
+ return loss_all
221
+
222
+
223
+ """
224
+ final train loss
225
+ """
226
+
227
+
228
+ def compute_dis_loss(
229
+ model,
230
+ flag_source_cat_single,
231
+ flag_source_cat_spatial,
232
+ anneal,
233
+ batch_features_all,
234
+ adj_all,
235
+ mask_batch_single,
236
+ mask_batch_spatial,
237
+ balance_weight_single_block,
238
+ balance_weight_spatial_block,
239
+ flagconfig,
240
+ ):
241
+ mask_batch_single_all = torch.hstack(mask_batch_single)
242
+ mask_batch_spatial_all = torch.hstack(mask_batch_spatial)
243
+ balance_weight_single_block = torch.hstack(balance_weight_single_block)
244
+ balance_weight_spatial_block = torch.hstack((balance_weight_spatial_block))
245
+
246
+ z_all = [
247
+ model.encoder["atlas" + str(i)](batch_features_all[i], adj_all[i])
248
+ for i in range(ModelType.n_atlas)
249
+ ]
250
+ z_mean_cat_single = torch.cat([z_all[i][3] for i in range(ModelType.n_atlas)])[
251
+ mask_batch_single_all, :
252
+ ]
253
+
254
+ z_spatial_all = [z_all[i][2] for i in range(ModelType.n_atlas)]
255
+ z_mean_cat_spatial = torch.cat(z_spatial_all)[mask_batch_spatial_all, :]
256
+
257
+ if anneal:
258
+ if z_mean_cat_single.shape[0] > 1:
259
+ noise_single = D.Normal(0, z_mean_cat_single.std(axis=0)).sample(
260
+ (z_mean_cat_single.shape[0],)
261
+ )
262
+ z_mean_cat_single = (
263
+ z_mean_cat_single
264
+ + (anneal * ModelType.align_noise_coef.value) * noise_single
265
+ )
266
+ if z_mean_cat_spatial.shape[0] > 1:
267
+ noise_spatial = D.Normal(
268
+ 0, ModelType.EPS.value + z_mean_cat_spatial.std(axis=0)
269
+ ).sample((z_mean_cat_spatial.shape[0],))
270
+ z_mean_cat_spatial = (
271
+ z_mean_cat_spatial
272
+ + (anneal * ModelType.align_noise_coef.value) * noise_spatial
273
+ )
274
+
275
+ ### compute dis loss
276
+ loss_dis_single = F.cross_entropy(
277
+ F.softmax(model.discriminator_single(z_mean_cat_single), dim=1),
278
+ flag_source_cat_single[mask_batch_single_all],
279
+ reduction="none",
280
+ )
281
+ loss_dis_single = (
282
+ balance_weight_single_block[mask_batch_single_all] * loss_dis_single
283
+ ).sum() / loss_dis_single.numel()
284
+
285
+ loss_dis_spatial = F.cross_entropy(
286
+ F.softmax(model.discriminator_spatial(z_mean_cat_spatial), dim=1),
287
+ flag_source_cat_spatial[mask_batch_spatial_all],
288
+ reduction="none",
289
+ )
290
+ loss_dis_spatial = (
291
+ balance_weight_spatial_block[mask_batch_spatial_all] * loss_dis_spatial
292
+ ).sum() / loss_dis_spatial.numel()
293
+
294
+ loss_dis = flagconfig.lambda_disc_single * (loss_dis_single + loss_dis_spatial)
295
+
296
+ loss_all = {"dis": loss_dis}
297
+ return loss_all
298
+
299
+
300
+ def compute_ae_loss(
301
+ model,
302
+ flag_source_cat_single,
303
+ flag_source_cat_spatial,
304
+ anneal,
305
+ batch_features_all,
306
+ adj_all,
307
+ mask_batch_single,
308
+ mask_batch_spatial,
309
+ balance_weight_single_block,
310
+ balance_weight_spatial_block,
311
+ flagconfig,
312
+ ):
313
+ z_all = [
314
+ model.encoder["atlas" + str(i)](batch_features_all[i], adj_all[i])
315
+ for i in range(ModelType.n_atlas)
316
+ ]
317
+
318
+ z_spatial_all = [z_all[i][2] for i in range(ModelType.n_atlas)]
319
+
320
+ decoder_all = [
321
+ model.decoder["atlas" + str(i)]( z_spatial_all[i], adj_all[i])
322
+ for i in range(ModelType.n_atlas)
323
+ ]
324
+
325
+ ### compute AE loss
326
+ # z_distribution_loss = [
327
+ # D.Normal(
328
+ # z_all[i][0][mask_batch_single[i], :], z_all[i][1][mask_batch_single[i], :]
329
+ # )
330
+ # for i in range(ModelType.n_atlas)
331
+ # ]
332
+ z_distribution_loss = [
333
+ z_all[i][0]
334
+ for i in range(ModelType.n_atlas)
335
+ ]
336
+ loss_AE_all = [
337
+ ModelType.lambda_ae_single.value
338
+ * AE_Gene_loss(
339
+ decoder_all[i][mask_batch_single[i], :],
340
+ batch_features_all[i][mask_batch_single[i], :],
341
+ z_distribution_loss[i],
342
+ )
343
+ for i in range(ModelType.n_atlas)
344
+ ]
345
+
346
+ mask_batch_single_all = torch.hstack(mask_batch_single)
347
+ mask_batch_spatial_all = torch.hstack(mask_batch_spatial)
348
+
349
+ z_mean_cat_single = torch.cat([z_all[i][3] for i in range(ModelType.n_atlas)])[
350
+ mask_batch_single_all, :
351
+ ]
352
+ z_mean_cat_spatial = torch.cat(z_spatial_all)[mask_batch_spatial_all, :]
353
+
354
+ if anneal:
355
+ if z_mean_cat_single.shape[0] > 1:
356
+ noise_single = D.Normal(0, z_mean_cat_single.std(axis=0)).sample(
357
+ (z_mean_cat_single.shape[0],)
358
+ )
359
+ z_mean_cat_single = (
360
+ z_mean_cat_single
361
+ + (anneal * ModelType.align_noise_coef.value) * noise_single
362
+ )
363
+ if z_mean_cat_spatial.shape[0] > 1:
364
+ noise_spatial = D.Normal(
365
+ 0, ModelType.EPS.value + z_mean_cat_spatial.std(axis=0)
366
+ ).sample((z_mean_cat_spatial.shape[0],))
367
+ z_mean_cat_spatial = (
368
+ z_mean_cat_spatial
369
+ + (anneal * ModelType.align_noise_coef.value) * noise_spatial
370
+ )
371
+
372
+ ### compute dis loss
373
+ balance_weight_single_block = torch.hstack(balance_weight_single_block)
374
+ balance_weight_spatial_block = torch.hstack((balance_weight_spatial_block))
375
+
376
+ loss_dis_single = F.cross_entropy(
377
+ F.softmax(model.discriminator_single(z_mean_cat_single), dim=1),
378
+ flag_source_cat_single[mask_batch_single_all],
379
+ reduction="none",
380
+ )
381
+ loss_dis_single = (
382
+ balance_weight_single_block[mask_batch_single_all] * loss_dis_single
383
+ ).sum() / loss_dis_single.numel()
384
+
385
+ loss_dis_spatial = F.cross_entropy(
386
+ F.softmax(model.discriminator_spatial(z_mean_cat_spatial), dim=1),
387
+ flag_source_cat_spatial[mask_batch_spatial_all],
388
+ reduction="none",
389
+ )
390
+ loss_dis_spatial = (
391
+ balance_weight_spatial_block[mask_batch_spatial_all] * loss_dis_spatial
392
+ ).sum() / loss_dis_spatial.numel()
393
+
394
+ loss_dis = flagconfig.lambda_disc_single * (loss_dis_single + loss_dis_spatial)
395
+
396
+ if (
397
+ flagconfig.lambda_disc_single == 1
398
+ ): # and loss_dis.item()<sum(loss_AE_all).item()/DIS_LAMDA:
399
+ flagconfig.lambda_disc_single = (
400
+ sum(loss_AE_all).item() / ModelType.DIS_LAMDA.value / loss_dis.item()
401
+ )
402
+ print(f"lambda_disc_single changed to {flagconfig.lambda_disc_single}")
403
+ loss_dis = flagconfig.lambda_disc_single * loss_dis
404
+
405
+ loss_all = {
406
+ "dis_ae": loss_dis,
407
+ "loss_AE_all": loss_AE_all,
408
+ "loss_all": -loss_dis + sum(loss_AE_all),
409
+ }
410
+ return loss_all
411
+
412
+
413
+ """
414
+ balance weight part
415
+ """
416
+
417
+
418
+ def get_balance_weight_subsample(leiden_adata_single, adatas_, key_leiden_category):
419
+ ########### function from GLUE: https://github.com/gao-lab/GLUE
420
+ us = [
421
+ sklearn.preprocessing.normalize(leiden.X, norm="l2")
422
+ for leiden in leiden_adata_single
423
+ ]
424
+ ns = [leiden.obs["size"] for leiden in leiden_adata_single]
425
+
426
+ power = 4
427
+ cutoff = 0.5
428
+ while True:
429
+ summary_balance_dict_sum = {}
430
+ summary_balance_dict_multiply = {}
431
+ summary_balance_dict_num = {}
432
+ for i, ui in enumerate(us):
433
+ for j, uj in enumerate(us[i + 1 :], start=i + 1):
434
+ cosine = ui @ uj.T
435
+ cosine[cosine < cutoff] = 0
436
+ cosine = COO.from_numpy(cosine)
437
+ cosine = np.power(cosine, power)
438
+
439
+ for ind in [i, j]:
440
+ if ind == i:
441
+ balancing = cosine.sum(axis=1).todense() / ns[ind]
442
+ else:
443
+ balancing = cosine.sum(axis=0).todense() / ns[ind]
444
+ balancing = pd.Series(
445
+ balancing, index=leiden_adata_single[ind].obs_names
446
+ )
447
+ balancing = balancing.loc[
448
+ adatas_[ind].obs[key_leiden_category]
449
+ ].to_numpy()
450
+ balancing /= balancing.sum() / balancing.size
451
+ if ind in summary_balance_dict_sum:
452
+ summary_balance_dict_sum[ind] += balancing.copy()
453
+ summary_balance_dict_multiply[ind] *= balancing.copy()
454
+ summary_balance_dict_num[ind] += 1
455
+ else:
456
+ summary_balance_dict_sum[ind] = balancing.copy()
457
+ summary_balance_dict_multiply[ind] = balancing.copy()
458
+ summary_balance_dict_num[ind] = 1
459
+ flag = 0
460
+ for i in range(len(summary_balance_dict_sum)):
461
+ if sum(np.isnan(summary_balance_dict_sum[i])) > 0:
462
+ flag = 1
463
+ break
464
+ for i in range(len(summary_balance_dict_multiply)):
465
+ if sum(np.isnan(summary_balance_dict_multiply[i])) > 0:
466
+ flag = 1
467
+ break
468
+ for i in range(len(summary_balance_dict_multiply)):
469
+ if sum(summary_balance_dict_sum[i]) == 0:
470
+ flag = 1
471
+ break
472
+ for i in range(len(summary_balance_dict_multiply)):
473
+ if sum(summary_balance_dict_multiply[i]) == 0:
474
+ flag = 1
475
+ break
476
+ if flag == 1:
477
+ cutoff -= 0.1
478
+ else:
479
+ break
480
+ print(f"balance weight final cutoff: {cutoff}")
481
+ for i in range(len(summary_balance_dict_sum)):
482
+ if (
483
+ summary_balance_dict_sum[i][summary_balance_dict_sum[i] == np.inf].shape[0]
484
+ > 0
485
+ ):
486
+ print(
487
+ i,
488
+ "inf:",
489
+ summary_balance_dict_sum[i][
490
+ summary_balance_dict_sum[i] == np.inf
491
+ ].shape[0],
492
+ )
493
+ summary_balance_dict_sum[i][summary_balance_dict_sum[i] == np.inf] = 1e308
494
+
495
+ for i in range(len(summary_balance_dict_sum)):
496
+ if (
497
+ summary_balance_dict_multiply[i][
498
+ summary_balance_dict_multiply[i] == np.inf
499
+ ].shape[0]
500
+ > 0
501
+ ):
502
+ print(
503
+ i,
504
+ "inf:",
505
+ summary_balance_dict_multiply[i][
506
+ summary_balance_dict_multiply[i] == np.inf
507
+ ].shape[0],
508
+ )
509
+ summary_balance_dict_multiply[i][
510
+ summary_balance_dict_multiply[i] == np.inf
511
+ ] = 1e308
512
+
513
+ balance_weight = []
514
+ summary_balance_dict = {}
515
+ for i in range(len(us)):
516
+ test1 = summary_balance_dict_sum[i] / (
517
+ summary_balance_dict_sum[i].sum() / summary_balance_dict_sum[i].size
518
+ )
519
+ test2 = summary_balance_dict_multiply[i] / (
520
+ summary_balance_dict_multiply[i].sum()
521
+ / summary_balance_dict_multiply[i].size
522
+ )
523
+ test = 0.9 * test1 + 0.1 * test2
524
+ test /= test.sum() / test.size
525
+ summary_balance_dict[i] = test.copy()
526
+ balance_weight.append(summary_balance_dict[i])
527
+ return balance_weight
528
+
529
+
530
+ def get_balance_weight(adatas, leiden_adata_single, adatas_, key_leiden_category):
531
+ ########### function from GLUE: https://github.com/gao-lab/GLUE
532
+ us = [
533
+ sklearn.preprocessing.normalize(leiden.X, norm="l2")
534
+ for leiden in leiden_adata_single
535
+ ]
536
+ ns = [leiden.obs["size"] for leiden in leiden_adata_single]
537
+
538
+ cosines = []
539
+ cutoff = 0.5
540
+ power = 4
541
+
542
+ for i, ui in enumerate(us):
543
+ for j, uj in enumerate(us[i + 1 :], start=i + 1):
544
+ cosine = ui @ uj.T
545
+ cosine[cosine < cutoff] = 0
546
+ cosine = COO.from_numpy(cosine)
547
+ cosine = np.power(cosine, power)
548
+ key = tuple(
549
+ slice(None) if k in (i, j) else np.newaxis for k in range(len(us))
550
+ ) # To align axes
551
+ cosines.append(cosine[key])
552
+ joint_cosine = prod(cosines)
553
+
554
+ if joint_cosine.coords.shape[0] == 0:
555
+ raise ValueError(
556
+ "Balance weight computation error! No correlation between samples or lower cutoff!"
557
+ )
558
+ #
559
+ balance_weight = []
560
+ for i, (adata, adata_, leiden, n) in enumerate(
561
+ zip(adatas, adatas_, leiden_adata_single, ns)
562
+ ):
563
+ balancing = (
564
+ joint_cosine.sum(
565
+ axis=tuple(k for k in range(joint_cosine.ndim) if k != i)
566
+ ).todense()
567
+ / n
568
+ )
569
+ balancing = pd.Series(balancing, index=leiden.obs_names)
570
+ balancing = balancing.loc[adata_.obs[key_leiden_category]].to_numpy()
571
+ balancing /= balancing.sum() / balancing.size
572
+ balance_weight.append(balancing)
573
+ return balance_weight
574
+
575
+
576
+ """
577
+ train ref data part
578
+ """
579
+
580
+
581
+ def compute_dis_loss_map(
582
+ adapt_model,
583
+ flag_source_cat_single,
584
+ flag_source_cat_spatial,
585
+ anneal,
586
+ batch_features_all,
587
+ adj_all,
588
+ mask_batch_single,
589
+ mask_batch_spatial,
590
+ pretrain_single_batch,
591
+ pretrain_spatial_batch,
592
+ flag_source_cat_single_pretrain,
593
+ flag_source_cat_spatial_pretrain,
594
+ flagconfig,
595
+ ):
596
+ mask_batch_single_all = torch.hstack(mask_batch_single)
597
+ mask_batch_spatial_all = torch.hstack(mask_batch_spatial)
598
+
599
+ z_all = [
600
+ adapt_model.encoder["atlas" + str(i)](batch_features_all[i], adj_all[i])
601
+ for i in range(ModelType.n_atlas)
602
+ ]
603
+ z_mean_cat_single = torch.cat([z_all[i][1] for i in range(ModelType.n_atlas)])[
604
+ mask_batch_single_all, :
605
+ ]
606
+ z_mean_cat_single = torch.vstack(
607
+ [
608
+ z_mean_cat_single,
609
+ torch.cat(
610
+ [pretrain_single_batch[i] for i in range(len(pretrain_single_batch))]
611
+ ),
612
+ ]
613
+ )
614
+
615
+ z_spatial_all = [z_all[i][2] for i in range(ModelType.n_atlas)]
616
+ z_mean_cat_spatial = torch.cat(z_spatial_all)[mask_batch_spatial_all, :]
617
+ z_mean_cat_spatial = torch.vstack(
618
+ [
619
+ z_mean_cat_spatial,
620
+ torch.cat(
621
+ [pretrain_spatial_batch[i] for i in range(len(pretrain_spatial_batch))]
622
+ ),
623
+ ]
624
+ )
625
+
626
+ ######### append pretrained data ##############
627
+
628
+ if anneal:
629
+ if z_mean_cat_single.shape[0] > 1:
630
+ noise_single = D.Normal(0, z_mean_cat_single.std(axis=0)).sample(
631
+ (z_mean_cat_single.shape[0],)
632
+ )
633
+ z_mean_cat_single = (
634
+ z_mean_cat_single
635
+ + (anneal * ModelType.align_noise_coef.value) * noise_single
636
+ )
637
+ if z_mean_cat_spatial.shape[0] > 1:
638
+ noise_spatial = D.Normal(
639
+ 0, ModelType.EPS.value + z_mean_cat_spatial.std(axis=0)
640
+ ).sample((z_mean_cat_spatial.shape[0],))
641
+ z_mean_cat_spatial = (
642
+ z_mean_cat_spatial
643
+ + (anneal * ModelType.align_noise_coef.value) * noise_spatial
644
+ )
645
+
646
+ ### compute dis loss
647
+ loss_dis_single = F.cross_entropy(
648
+ F.softmax(
649
+ torch.hstack(
650
+ [
651
+ adapt_model.discriminator_single(z_mean_cat_single),
652
+ adapt_model.discriminator_single_pretrain(z_mean_cat_single),
653
+ ]
654
+ ),
655
+ dim=1,
656
+ ),
657
+ torch.hstack(
658
+ [
659
+ flag_source_cat_single[mask_batch_single_all],
660
+ flag_source_cat_single_pretrain,
661
+ ]
662
+ ),
663
+ reduction="none",
664
+ )
665
+ loss_dis_single = loss_dis_single.sum() / loss_dis_single.numel()
666
+
667
+ loss_dis_spatial = F.cross_entropy(
668
+ F.softmax(
669
+ torch.hstack(
670
+ [
671
+ adapt_model.discriminator_spatial(z_mean_cat_spatial),
672
+ adapt_model.discriminator_spatial_pretrain(z_mean_cat_spatial),
673
+ ]
674
+ ),
675
+ dim=1,
676
+ ),
677
+ torch.hstack(
678
+ [
679
+ flag_source_cat_spatial[mask_batch_spatial_all],
680
+ flag_source_cat_spatial_pretrain,
681
+ ]
682
+ ),
683
+ reduction="none",
684
+ )
685
+ loss_dis_spatial = loss_dis_spatial.sum() / loss_dis_spatial.numel()
686
+
687
+ loss_dis = flagconfig.lambda_disc_single * (loss_dis_single + loss_dis_spatial)
688
+ # loss_dis = self.lambda_disc_single * (loss_dis_single )
689
+
690
+ loss_all = {"dis": loss_dis}
691
+ return loss_all
692
+
693
+
694
+ def compute_ae_loss_map(
695
+ adapt_model,
696
+ flag_source_cat_single,
697
+ flag_source_cat_spatial,
698
+ anneal,
699
+ batch_features_all,
700
+ adj_all,
701
+ mask_batch_single,
702
+ mask_batch_spatial,
703
+ pretrain_single_batch,
704
+ pretrain_spatial_batch,
705
+ flag_source_cat_single_pretrain,
706
+ flag_source_cat_spatial_pretrain,
707
+ flagconfig
708
+ ):
709
+ z_all = [
710
+ adapt_model.encoder["atlas" + str(i)](batch_features_all[i], adj_all[i])
711
+ for i in range(ModelType.n_atlas)
712
+ ]
713
+
714
+ # z_distribution_all = [
715
+ # z_all[i][0] for i in range(ModelType.n_atlas)
716
+ # ]
717
+ # z_sample_all = [z_distribution_all[i].rsample() for i in range(ModelType.n_atlas)]
718
+
719
+ z_spatial_all = [z_all[i][2] for i in range(ModelType.n_atlas)]
720
+
721
+ decoder_all = [
722
+ adapt_model.decoder["atlas" + str(i)](
723
+ z_all[i][1],
724
+ z_spatial_all[i],
725
+ adj_all[i],
726
+ adapt_model.gene_embedding_pretrained,
727
+ adapt_model.gene_embedding_new,
728
+ )
729
+ for i in range(ModelType.n_atlas)
730
+ ]
731
+
732
+ ### compute AE loss
733
+ z_distribution_loss = [
734
+ z_all[i][0]
735
+ for i in range(ModelType.n_atlas)
736
+ ]
737
+ loss_AE_all = [
738
+ ModelType.lambda_ae_single.value
739
+ * AE_Gene_loss(
740
+ decoder_all[i][mask_batch_single[i], :],
741
+ batch_features_all[i][mask_batch_single[i], :],
742
+ z_distribution_loss[i],
743
+ )
744
+ for i in range(ModelType.n_atlas)
745
+ ]
746
+
747
+ mask_batch_single_all = torch.hstack(mask_batch_single)
748
+ mask_batch_spatial_all = torch.hstack(mask_batch_spatial)
749
+
750
+ z_mean_cat_single = torch.cat([z_all[i][1] for i in range(ModelType.n_atlas)])[
751
+ mask_batch_single_all, :
752
+ ]
753
+ z_mean_cat_single = torch.vstack(
754
+ [
755
+ z_mean_cat_single,
756
+ torch.cat(
757
+ [pretrain_single_batch[i] for i in range(len(pretrain_single_batch))]
758
+ ),
759
+ ]
760
+ )
761
+
762
+ z_mean_cat_spatial = torch.cat(z_spatial_all)[mask_batch_spatial_all, :]
763
+ z_mean_cat_spatial = torch.vstack(
764
+ [
765
+ z_mean_cat_spatial,
766
+ torch.cat(
767
+ [pretrain_spatial_batch[i] for i in range(len(pretrain_spatial_batch))]
768
+ ),
769
+ ]
770
+ )
771
+
772
+ if anneal:
773
+ if z_mean_cat_single.shape[0] > 1:
774
+ noise_single = D.Normal(0, z_mean_cat_single.std(axis=0)).sample(
775
+ (z_mean_cat_single.shape[0],)
776
+ )
777
+ z_mean_cat_single = (
778
+ z_mean_cat_single + (anneal * ModelType.align_noise_coef.value) * noise_single
779
+ )
780
+ if z_mean_cat_spatial.shape[0] > 1:
781
+ noise_spatial = D.Normal(0, ModelType.EPS.value + z_mean_cat_spatial.std(axis=0)).sample(
782
+ (z_mean_cat_spatial.shape[0],)
783
+ )
784
+ z_mean_cat_spatial = (
785
+ z_mean_cat_spatial + (anneal * ModelType.align_noise_coef.value) * noise_spatial
786
+ )
787
+
788
+ ### compute dis loss
789
+ loss_dis_single = F.cross_entropy(
790
+ F.softmax(
791
+ torch.hstack(
792
+ [
793
+ adapt_model.discriminator_single(z_mean_cat_single),
794
+ adapt_model.discriminator_single_pretrain(z_mean_cat_single),
795
+ ]
796
+ ),
797
+ dim=1,
798
+ ),
799
+ torch.hstack(
800
+ [
801
+ flag_source_cat_single[mask_batch_single_all],
802
+ flag_source_cat_single_pretrain,
803
+ ]
804
+ ),
805
+ reduction="none",
806
+ )
807
+ loss_dis_single = loss_dis_single.sum() / loss_dis_single.numel()
808
+
809
+ loss_dis_spatial = F.cross_entropy(
810
+ F.softmax(
811
+ torch.hstack(
812
+ [
813
+ adapt_model.discriminator_spatial(z_mean_cat_spatial),
814
+ adapt_model.discriminator_spatial_pretrain(z_mean_cat_spatial),
815
+ ]
816
+ ),
817
+ dim=1,
818
+ ),
819
+ torch.hstack(
820
+ [
821
+ flag_source_cat_spatial[mask_batch_spatial_all],
822
+ flag_source_cat_spatial_pretrain,
823
+ ]
824
+ ),
825
+ reduction="none",
826
+ )
827
+ loss_dis_spatial = loss_dis_spatial.sum() / loss_dis_spatial.numel()
828
+
829
+ loss_dis = flagconfig.lambda_disc_single * (loss_dis_single + loss_dis_spatial)
830
+
831
+ if (
832
+ flagconfig.lambda_disc_single == 1
833
+ ): # and loss_dis.item()<sum(loss_AE_all).item()/DIS_LAMDA:
834
+ flagconfig.lambda_disc_single = sum(loss_AE_all).item() / ModelType.DIS_LAMDA.value / loss_dis.item()
835
+ logging.info(f"\n\nlambda_disc_single changed to {flagconfig.lambda_disc_single}\n")
836
+ loss_dis = flagconfig.lambda_disc_single * loss_dis
837
+
838
+ loss_all = {
839
+ "dis_ae": loss_dis,
840
+ "loss_AE_all": loss_AE_all,
841
+ "loss_all": -loss_dis + sum(loss_AE_all),
842
+ }
843
+ return loss_all