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
« prev ^ index » next coverage.py v7.11.0, created at 2026-07-23 14:40 +0000
1"""
22-level DAG generation
3"""
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
11class GenDAG2Level:
12 """
13 generate a DAG with 2 levels: first level generate macro nodes, second
14 level populate each macro node
15 """
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
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
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)
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)
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()
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)
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 )
82 self.dag_refined.add_arc_ind(ind_global_tail, ind_global_head)
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)
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)
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))
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 ]
114 def run(self):
115 """
116 generation
117 """
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