-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathexample_simulation.py
More file actions
151 lines (117 loc) · 4.07 KB
/
Copy pathexample_simulation.py
File metadata and controls
151 lines (117 loc) · 4.07 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
"""
Synthetic example for exact and greedy local-BIC refinement.
This example generates a true DAG, simulates linear-Gaussian data,
adds false-positive edges to create an initial estimated DAG, and then
uses exact local subset selection and greedy deletion to remove unsupported
edges from the same fixed candidate DAG.
Run:
python example_simulation.py
"""
import numpy as np
from local_bic_refinement import (
exact_refine_dag,
graph_metrics,
greedy_refine_dag,
is_acyclic,
)
def random_dag(d: int, expected_edges: int, rng: np.random.Generator) -> np.ndarray:
"""Generate a random DAG by sampling edges in a random topological order."""
order = rng.permutation(d)
A = np.zeros((d, d), dtype=int)
p = expected_edges / (d * (d - 1) / 2)
for a in range(d):
for b in range(a + 1, d):
if rng.random() < p:
parent = order[a]
child = order[b]
A[parent, child] = 1
return A
def topological_order(A: np.ndarray) -> list[int]:
"""Return one topological order for a DAG."""
A = np.asarray(A, dtype=int)
indegree = A.sum(axis=0).astype(int)
queue = list(np.where(indegree == 0)[0])
order = []
while queue:
node = queue.pop(0)
order.append(int(node))
for child in np.where(A[node, :] != 0)[0]:
indegree[child] -= 1
if indegree[child] == 0:
queue.append(int(child))
if len(order) != A.shape[0]:
raise ValueError("A is not acyclic.")
return order
def simulate_linear_sem(
A: np.ndarray,
n: int,
rng: np.random.Generator,
weight_low: float = 0.5,
weight_high: float = 2.0,
) -> np.ndarray:
"""Simulate data from a linear SEM using a DAG adjacency matrix."""
d = A.shape[0]
W = np.zeros((d, d), dtype=float)
for i, j in zip(*np.where(A != 0)):
sign = rng.choice([-1.0, 1.0])
W[i, j] = sign * rng.uniform(weight_low, weight_high)
X = np.zeros((n, d), dtype=float)
noise = rng.normal(size=(n, d))
for j in topological_order(A):
parents = np.where(A[:, j] != 0)[0]
X[:, j] = noise[:, j]
if len(parents) > 0:
X[:, j] += X[:, parents] @ W[parents, j]
return X
def add_false_positive_edges(
A: np.ndarray,
n_extra: int,
rng: np.random.Generator,
) -> np.ndarray:
"""
Add false-positive edges while preserving acyclicity.
This mimics an initial DAG estimate that contains extra edges.
"""
d = A.shape[0]
A_initial = A.copy()
candidates = [(i, j) for i in range(d) for j in range(d) if i != j and A_initial[i, j] == 0]
rng.shuffle(candidates)
added = 0
for i, j in candidates:
A_initial[i, j] = 1
if is_acyclic(A_initial):
added += 1
else:
A_initial[i, j] = 0
if added >= n_extra:
break
return A_initial
def main() -> None:
rng = np.random.default_rng(123)
d = 10
n = 500
expected_edges = 2 * d
A_true = random_dag(d=d, expected_edges=expected_edges, rng=rng)
X = simulate_linear_sem(A_true, n=n, rng=rng)
A_initial = add_false_positive_edges(A_true, n_extra=10, rng=rng)
exact = exact_refine_dag(X, A_initial)
greedy = greedy_refine_dag(X, A_initial)
print("Fixed-candidate local-BIC refinement example")
print("--------------------------------------------")
print(f"True edges: {int(A_true.sum())}")
print(f"Candidate edges: {int(A_initial.sum())}")
print(f"Exact edges: {int(exact.adjacency.sum())}")
print(f"Greedy edges: {int(greedy.adjacency.sum())}")
print(f"Exact certified: {exact.globally_optimal}")
print(f"Greedy BIC gap: {greedy.total_bic - exact.total_bic:.6f}")
print()
print("Initial graph diagnostics")
print(graph_metrics(A_true, A_initial))
print()
print("Exact-refinement diagnostics")
print(graph_metrics(A_true, exact.adjacency))
print()
print("Greedy-refinement diagnostics")
print(graph_metrics(A_true, greedy.adjacency))
if __name__ == "__main__":
main()