ProCreations's picture
Publish validated ICML reproduction
177308a verified
Raw
History Blame Contribute Delete
7.18 kB
"""
Finite-sample counterparts of the graphical results, using the same causal-learn
implementations the paper used.
(A) Lemma 1 in data: a selection-blind skeleton search on evolutionary data
produces adjacencies that the static selection model says cannot be there;
the same search on data WITHOUT selection (pa(S) = {}) must not.
(B) Theorem 2 in data: PC on evolutionary data, oriented vs all adjacencies.
(C) Theorem 4 in data: CDNOD on K = 4 heterogeneous domains (selection
mechanism changed) vs PC on a single domain.
Run: python3 finite_sample.py
"""
import json, os, warnings, itertools, time
import numpy as np
warnings.filterwarnings('ignore')
os.environ.setdefault('OMP_NUM_THREADS', '1')
OUT = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'outputs')
os.makedirs(OUT, exist_ok=True)
N = 5000
ALPHA = 0.05
def skeleton(X, alpha=ALPHA):
from causallearn.search.ConstraintBased.PC import pc
from exp6_synthetic import _edges_from_cl
cg = pc(X, alpha=alpha, indep_test='fisherz', show_progress=False)
dd, uu = _edges_from_cl(cg.G, X.shape[1])
return dd, uu, set(frozenset(e) for e in dd) | uu
def task_A(seed):
"""Lemma 1 in data."""
from evosel import (random_static_dag, simulate_evolution, clique_augmented,
SEL, DG)
d = 8
rng = np.random.default_rng(seed)
G = random_static_dag(d, rng, avg_deg=2.0, n_sel_parents=3)
true = set(frozenset((i, j)) for j in range(d) for i in G.pa[j] if i != SEL)
gp = clique_augmented(G, d)
gpe = set(frozenset(e) for e in gp.edges())
B, w, var = None, None, None
Xe, p = simulate_evolution(G, d, 3, N, rng) # with selection
Xn, _ = simulate_evolution(G, d, 3, N, rng, selection=False,
B=p[0], w=p[1], var=p[2]) # reproduction at random
_, _, ske = skeleton(Xe)
_, _, skn = skeleton(Xn)
return {'seed': seed, 'd': d,
'n_true_causal_edges': len(true),
'evolution_adjacencies': len(ske),
'evolution_spurious_wrt_true_causal': len(ske - true),
'evolution_spurious_explained_by_Gplus': len((ske - true) & gpe),
'no_selection_adjacencies': len(skn),
'no_selection_spurious_wrt_true_causal': len(skn - true)}
def task_B(seed):
"""Theorem 2 in data (d = 10, the smallest size the paper reports)."""
from evosel import random_static_dag, simulate_evolution, clique_augmented, cpdag, SEL
from exp6_synthetic import _score
d = 10
rng = np.random.default_rng(1000 + seed)
G = random_static_dag(d, rng, avg_deg=2.0)
true = set((i, j) for j in range(d) for i in G.pa[j] if i != SEL)
anS = G.ancestors([SEL]) - {SEL}
X, _ = simulate_evolution(G, d, 3, N, rng)
dd, uu, _ = skeleton(X)
s = _score(dd, uu, true, anS)
s.update({'seed': seed, 'd': d, 'n_anS': len(anS)})
return s
def task_C(seed):
"""Theorem 4 in data: CDNOD over 4 domains where the selection mechanism
changes, versus PC on domain 1 alone."""
from causallearn.search.ConstraintBased.CDNOD import cdnod
from evosel import (random_static_dag, simulate_evolution, sem_params,
clique_augmented, cpdag, multidomain_augmented, SEL)
from exp6_synthetic import _edges_from_cl, _score
d, K = 8, 4
rng = np.random.default_rng(2000 + seed)
G = random_static_dag(d, rng, avg_deg=2.0, n_sel_parents=2)
true = set((i, j) for j in range(d) for i in G.pa[j] if i != SEL)
anS = G.ancestors([SEL]) - {SEL}
B, w, var = sem_params(G, d, rng)
Xs = []
for k in range(K): # only the selection weights change
wk = w * rng.uniform(0.3, 2.0, size=d) if k else w
Xk, _ = simulate_evolution(G, d, 3, N, rng, B=B, w=wk, var=var)
Xs.append(Xk)
X = np.vstack(Xs)
cidx = np.repeat(np.arange(K), N).reshape(-1, 1)
cg = cdnod(X, cidx, alpha=ALPHA, indep_test='fisherz', show_progress=False)
dd, uu = _edges_from_cl(cg.G, d + 1)
dd = set((u, v) for (u, v) in dd if u < d and v < d)
uu = set(e for e in uu if all(x < d for x in e))
multi = _score(dd, uu, true, anS)
d1, u1, _ = skeleton(Xs[0])
single = _score(d1, u1, true, anS)
return {'seed': seed, 'd': d, 'K': K, 'n_anS': len(anS),
'single_domain': single, 'multi_domain': multi,
'single_correct_oriented': single['oriented_direction_correct'],
'multi_correct_oriented': multi['oriented_direction_correct'],
'single_n_oriented': single['n_oriented'],
'multi_n_oriented': multi['n_oriented']}
def mean(rows, path):
v = []
for r in rows:
x = r
for p in path:
x = x[p]
v.append(x)
a = np.array(v, dtype=float)
a = a[~np.isnan(a)]
return float(a.mean())
if __name__ == '__main__':
import multiprocessing as mp
t0 = time.time()
with mp.Pool(processes=min(mp.cpu_count(), 12)) as pool:
A = pool.map(task_A, range(20))
B = pool.map(task_B, range(20))
C = pool.map(task_C, range(20))
res = {
'A_lemma1_in_data': {
'runs': len(A), 'd': 8, 'N': N, 'T': 3,
'mean_adjacencies_with_evolution': mean(A, ['evolution_adjacencies']),
'total_spurious_with_evolution': sum(r['evolution_spurious_wrt_true_causal'] for r in A),
'total_spurious_explained_by_Gplus': sum(r['evolution_spurious_explained_by_Gplus'] for r in A),
'total_spurious_without_selection': sum(r['no_selection_spurious_wrt_true_causal'] for r in A),
'mean_spurious_with_evolution': mean(A, ['evolution_spurious_wrt_true_causal']),
'mean_spurious_without_selection': mean(A, ['no_selection_spurious_wrt_true_causal']),
'detail': A},
'B_theorem2_in_data': {
'runs': len(B), 'd': 10, 'N': N, 'T': 3,
'precision_standard': mean(B, ['precision_standard']),
'precision_ours': mean(B, ['precision_ours']),
'theorem2_soundness_rate': (
sum(r['oriented_satisfying_theorem2'] for r in B) /
max(sum(r['n_oriented'] for r in B), 1)),
'runs_where_ours_ge_standard': sum(
1 for r in B if r['precision_ours'] >= r['precision_standard']),
'detail': B},
'C_theorem4_in_data': {
'runs': len(C), 'd': 8, 'K': 4, 'N_per_domain': N,
'mean_correct_oriented_single': mean(C, ['single_correct_oriented']),
'mean_correct_oriented_multi': mean(C, ['multi_correct_oriented']),
'mean_n_oriented_single': mean(C, ['single_n_oriented']),
'mean_n_oriented_multi': mean(C, ['multi_n_oriented']),
'runs_multi_ge_single': sum(
1 for r in C if r['multi_correct_oriented'] >= r['single_correct_oriented']),
'detail': C},
'runtime_sec': time.time() - t0}
json.dump(res, open(os.path.join(OUT, 'finite_sample.json'), 'w'), indent=1)
for k in ('A_lemma1_in_data', 'B_theorem2_in_data', 'C_theorem4_in_data'):
print(k, {a: b for a, b in res[k].items() if a != 'detail'})