Coverage for src/causalspyne/gen_dag_2level.py: 90%

61 statements  

« prev     ^ index     » next       coverage.py v7.11.0, created at 2026-07-23 14:40 +0000

1""" 

22-level DAG generation 

3""" 

4 

5from causalspyne.dag_stack_indexer import DAGStackIndexer 

6from causalspyne.dag_manipulator import DAGManipulator 

7from causalspyne.weight import WeightGenWishart 

8from causalspyne.utils_random import coerce_rng 

9 

10 

11class GenDAG2Level: 

12 """ 

13 generate a DAG with 2 levels: first level generate macro nodes, second 

14 level populate each macro node 

15 """ 

16 

17 def __init__( 

18 self, 

19 dag_generator, 

20 num_macro_nodes, 

21 num_micro_nodes, 

22 max_num_local_nodes=4, 

23 min_num_local_nodes=3, 

24 rng=None, 

25 ): 

26 rng = coerce_rng(rng) 

27 self.dag_generator = dag_generator 

28 self.num_macro_nodes = num_macro_nodes 

29 self.num_micro_nodes = num_micro_nodes 

30 self.max_num_local_nodes = max_num_local_nodes 

31 self.min_num_local_nodes = min_num_local_nodes 

32 

33 self.global_dag_indexer = None 

34 self.dag_backbone = None 

35 self.dict_macro_node2dag = {} 

36 self.dag_refined = None 

37 self.rng = rng 

38 self.dag_manipulator = None 

39 

40 def populate_macro_node(self): 

41 """ 

42 replace a macro node into a DAG 

43 """ 

44 # iterate each macro node 

45 for name in self.dag_backbone.list_node_names: 

46 num_nodes = self.num_micro_nodes 

47 if num_nodes is None: 

48 num_nodes = self.rng.integers(self.min_num_local_nodes, self.max_num_local_nodes + 1) 

49 self.dict_macro_node2dag[name] = self.dag_generator.gen_dag( 

50 num_nodes=num_nodes, 

51 prefix=name, 

52 target_num_confounder=2, 

53 ) 

54 self.global_dag_indexer = DAGStackIndexer(self) 

55 

56 def interconnection(self): 

57 """ 

58 connect macro nodes with edges 

59 """ 

60 # iterate over the Macro-DAG edges 

61 for arc in self.dag_backbone.list_arcs: 

62 self.connect_macro_node_via_local_node(arc) 

63 

64 def connect_macro_node_via_local_node(self, arc): 

65 """ 

66 connect macro-DAG node edge (i,j) via local nodes 

67 """ 

68 macro_arrow_tail, macro_arrow_head = arc 

69 _, ind_local_tail = self.dict_macro_node2dag[macro_arrow_tail].sample_node() 

70 _, ind_local_head = self.dict_macro_node2dag[macro_arrow_head].sample_node() 

71 

72 ind_macro_tail = self.dag_backbone.get_node_ind(macro_arrow_tail) 

73 ind_macro_head = self.dag_backbone.get_node_ind(macro_arrow_head) 

74 

75 ind_global_tail = self.global_dag_indexer.get_global_ind( 

76 ind_macro_tail, ind_local_tail 

77 ) 

78 ind_global_head = self.global_dag_indexer.get_global_ind( 

79 ind_macro_head, ind_local_head 

80 ) 

81 

82 self.dag_refined.add_arc_ind(ind_global_tail, ind_global_head) 

83 

84 def inject_additional_confounder(self): 

85 """ 

86 make confounder in the big graph 

87 """ 

88 obj_gen_weight = WeightGenWishart(rng=self.rng) 

89 self.dag_manipulator = DAGManipulator(self.dag_refined, 

90 obj_gen_weight, self.rng) 

91 

92 ind_arbitrary = self.dag_refined.get_top_last() 

93 self.dag_manipulator.mk_confound(ind_arbitrary) 

94 print(self.dag_refined.num_confounder) 

95 ind_arbitrary = self.dag_refined.climb(ind_arbitrary) 

96 self.dag_manipulator.mk_confound(ind_arbitrary) 

97 print(self.dag_refined.num_confounder) 

98 

99 def get_macro_node_global_inds(self, macro_name): 

100 """Return the list of global micro-node indices for a given macro node name.""" 

101 ind_macro = list(self.dag_backbone.list_node_names).index(macro_name) 

102 num_micro = self.global_dag_indexer.dict_num[macro_name] 

103 start = self.global_dag_indexer.list_accum_count[ind_macro] 

104 return list(range(start, start + num_micro)) 

105 

106 def get_root_macro_names(self): 

107 """Return names of macro nodes that have no incoming edges (roots).""" 

108 mat = self.dag_backbone.mat_adjacency 

109 return [ 

110 name for i, name in enumerate(self.dag_backbone.list_node_names) 

111 if mat[i, :].sum() == 0 

112 ] 

113 

114 def run(self): 

115 """ 

116 generation 

117 """ 

118 

119 # generate dag_backbone DAG with only macro nodes 

120 self.dag_backbone = self.dag_generator.gen_dag( 

121 self.num_macro_nodes, target_num_confounder=2 

122 ) 

123 self.populate_macro_node() 

124 self.interconnection() 

125 self.dag_refined.check() 

126 self.inject_additional_confounder() 

127 return self.dag_refined