11import torch
22from torch_geometric .data import Batch , Data
33
4+ from decipher import CFG
45from decipher .explain .gene .gene_selection import train_GAE
56from decipher .explain .regress .mixin import train_regress
6- from decipher .utils import GENESELECT_CFG , REGRESS_CFG
77
88
99def test_GAE_mimic ():
10- GENESELECT_CFG .center_dim = 8
11- GENESELECT_CFG .expr_dim = 100
12- GENESELECT_CFG .work_dir = "./results/explain"
13- GENESELECT_CFG .gae_epochs = 2
14- GENESELECT_CFG .fit .epochs = 2
10+ CFG .gene_select .center_dim = 8
11+ CFG .gene_select .expr_dim = 100
12+ CFG .gene_select .work_dir = "./results/explain"
13+ CFG .gene_select .gae_epochs = 2
1514
1615 N_CELL1 = 100
1716 graph1 = Data (
@@ -27,19 +26,19 @@ def test_GAE_mimic():
2726 expr = torch .randn (N_CELL2 , 100 ),
2827 )
2928 graph_all = Batch .from_data_list ([graph1 , graph2 ])
30- train_GAE (graph1 , GENESELECT_CFG , save_dir = "test_GAE" )
31- train_GAE (graph_all , GENESELECT_CFG , save_dir = "test_GAE_batched" )
29+ train_GAE (graph1 , CFG . gene_select , save_dir = "test_GAE" )
30+ train_GAE (graph_all , CFG . gene_select , save_dir = "test_GAE_batched" )
3231
3332
3433def test_regress_mimic ():
35- REGRESS_CFG .center_dim = 8
36- REGRESS_CFG .nbr_dim = 8
37- REGRESS_CFG .work_dir = "./results/explain"
38- REGRESS_CFG . fit .epochs = 2
34+ CFG . regress .center_dim = 8
35+ CFG . regress .nbr_dim = 8
36+ CFG . regress .work_dir = "./results/explain"
37+ CFG . regress . trainer .epochs = 2
3938
4039 x = torch .randn (100 , 8 )
4140 y = torch .randn (100 , 8 )
42- train_regress (x , y , REGRESS_CFG , save_dir = "test_regress" )
41+ train_regress (x , y , CFG . regress , save_dir = "test_regress" )
4342
4443
4544if __name__ == "__main__" :
0 commit comments