""" CORRECTNESS GATES. Everything downstream depends on two primitives: (a) the d-separation oracle -> cross-checked against networkx (b) the DAG -> CPDAG routine (Meek) -> cross-checked against causal-learn plus two re-derivations of results the literature already establishes: (c) Verma-Pearl: the CPDAG's directed edges are exactly the edges that are invariant over the Markov equivalence class (checked by brute-force enumeration of every DAG with the same skeleton & v-structures); (d) the global Markov property of the simulated linear-Gaussian SCM (d-separation <=> vanishing partial correlation), which gates the finite-sample experiments. Run: python3 gates.py """ import itertools, json, os, warnings import numpy as np import networkx as nx from evosel import (DG, SEL, dsep, cpdag, random_static_dag, clique_augmented, evolutionary_graph, simulate_evolution, _traits, sem_params) warnings.filterwarnings('ignore') OUT = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'outputs') os.makedirs(OUT, exist_ok=True) res = {} def to_nx(g): G = nx.DiGraph() G.add_nodes_from(g.nodes) G.add_edges_from(g.edges()) return G # --------------------------------------------------------------- (a) d-sep def gate_dsep(n_models=400, seed=0): rng = np.random.default_rng(seed) n_q = 0 bad = 0 for _ in range(n_models): d = int(rng.integers(3, 8)) G = random_static_dag(d, rng, avg_deg=float(rng.uniform(1.0, 3.0))) T = int(rng.integers(1, 4)) g = evolutionary_graph(G, d, T) nxg = to_nx(g) nodes = list(g.nodes) for _ in range(30): k = rng.integers(0, 3, size=len(nodes)) A = [v for v, kk in zip(nodes, k) if kk == 0] B = [v for v, kk in zip(nodes, k) if kk == 1] C = [v for v, kk in zip(nodes, k) if kk == 2] if not A or not B: continue mine = dsep(g, A, B, C) theirs = nx.is_d_separator(nxg, set(A), set(B), set(C)) n_q += 1 bad += (mine != theirs) return {'queries': n_q, 'disagreements': bad} # --------------------------------------------------------------- (b) CPDAG def gate_cpdag(n_models=600, seed=1): from causallearn.graph.Dag import Dag from causallearn.graph.GraphNode import GraphNode from causallearn.utils.DAG2CPDAG import dag2cpdag from causallearn.graph.Endpoint import Endpoint rng = np.random.default_rng(seed) bad = 0 n = 0 for _ in range(n_models): d = int(rng.integers(3, 9)) G = random_static_dag(d, rng, avg_deg=float(rng.uniform(1.0, 3.5))) gp = clique_augmented(G, d) mine_d, mine_u = cpdag(gp) # causal-learn oracle nds = [GraphNode('V%d' % i) for i in range(d)] dag = Dag(nds) for (u, v) in gp.edges(): dag.add_directed_edge(nds[u], nds[v]) cp = dag2cpdag(dag) cl_d, cl_u = set(), set() for e in cp.get_graph_edges(): i = int(e.get_node1().get_name()[1:]) j = int(e.get_node2().get_name()[1:]) e1, e2 = e.get_endpoint1(), e.get_endpoint2() if e1 == Endpoint.TAIL and e2 == Endpoint.ARROW: cl_d.add((i, j)) elif e1 == Endpoint.ARROW and e2 == Endpoint.TAIL: cl_d.add((j, i)) else: cl_u.add(frozenset((i, j))) n += 1 bad += (mine_d != cl_d or mine_u != cl_u) return {'models': n, 'disagreements': bad} # ------------------------------------------- (c) Verma-Pearl invariance check def all_dags(d): """All labelled DAGs on d nodes (exhaustive).""" pairs = list(itertools.combinations(range(d), 2)) out = [] for code in itertools.product((0, 1, 2), repeat=len(pairs)): g = DG(range(d)) for (a, b), c in zip(pairs, code): if c == 1: g.add(a, b) elif c == 2: g.add(b, a) if g.is_acyclic(): out.append(g) return out def gate_verma(d=4): dags = all_dags(d) groups = {} for g in dags: dd, uu = cpdag(g) key = (tuple(sorted(dd)), tuple(sorted(tuple(sorted(e)) for e in uu))) groups.setdefault(key, []).append(g) bad_dir = bad_und = 0 for key, members in groups.items(): dd = set(key[0]) # every member must contain every compelled edge with the same direction for g in members: es = set(g.edges()) if not dd <= es: bad_dir += 1 # every undirected edge must be reversible somewhere in the class for e in key[1]: a, b = e if not (any(g.has(a, b) for g in members) and any(g.has(b, a) for g in members)): bad_und += 1 return {'n_dags': len(dags), 'n_equivalence_classes': len(groups), 'compelled_edge_violations': bad_dir, 'reversibility_violations': bad_und} # --------------------------------------- (d) global Markov property of the SCM def partial_corr(Sig, i, j, C): idx = [i, j] + list(C) M = np.linalg.inv(Sig[np.ix_(idx, idx)]) return -M[0, 1] / np.sqrt(M[0, 0] * M[1, 1]) def gate_markov(seeds=(7, 8, 9), n=400000, d=6): """Population-level check on the *evolutionary* DGP: partial correlations must vanish for the d-separations of G^+ (Theorem 1) and not otherwise. Also a directly falsifiable control: triples that the selection-blind static graph (G with S deleted) declares independent but G^+ does not.""" out = [] for seed in seeds: rng = np.random.default_rng(seed) G = random_static_dag(d, rng, avg_deg=2.0, n_sel_parents=2) X, _ = simulate_evolution(G, d, T=3, n=n, rng=rng) Sig = np.cov(X.T) gp = clique_augmented(G, d) naive = DG(range(d)) for j in range(d): for i in G.pa[j]: if i != SEL: naive.add(i, j) sep, con, ctrl = [], [], [] for i, j in itertools.combinations(range(d), 2): rest = [k for k in range(d) if k not in (i, j)] for r in range(len(rest) + 1): for C in itertools.combinations(rest, r): pc = abs(partial_corr(Sig, i, j, C)) s_plus = dsep(gp, [i], [j], list(C)) (sep if s_plus else con).append(pc) if (not s_plus) and dsep(naive, [i], [j], list(C)): ctrl.append(pc) mx = float(max(sep)) if sep else 0.0 out.append({ 'seed': seed, 'n_dsep_triples': len(sep), 'n_dconn_triples': len(con), 'max_abs_pcorr_when_dseparated_in_Gplus': mx, 'frac_dconnected_above_that_max': float(np.mean(np.array(con) > mx)), 'median_abs_pcorr_when_dconnected': float(np.median(con)), 'n_control_triples_static_says_indep': len(ctrl), 'median_abs_pcorr_on_control_triples': float(np.median(ctrl)) if ctrl else None, 'max_abs_pcorr_on_control_triples': float(max(ctrl)) if ctrl else None}) return {'n_samples': int(n), 'd': d, 'T': 3, 'per_seed': out} if __name__ == '__main__': res['gate_a_dseparation_vs_networkx'] = gate_dsep() print('a', res['gate_a_dseparation_vs_networkx']) res['gate_b_cpdag_vs_causallearn'] = gate_cpdag() print('b', res['gate_b_cpdag_vs_causallearn']) res['gate_c_verma_pearl_invariance'] = gate_verma() print('c', res['gate_c_verma_pearl_invariance']) res['gate_d_global_markov_of_simulated_scm'] = gate_markov() print('d', res['gate_d_global_markov_of_simulated_scm']) json.dump(res, open(os.path.join(OUT, 'gates.json'), 'w'), indent=1) print('written')