summaryrefslogtreecommitdiff
path: root/experiments/analyze_c2_nodepert_validation.py
blob: 634ea0da6fcd3c9ba1f15bd9ae6a13ec6caca267 (plain)
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
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
"""Audit the frozen C2 direct node-perturbation causal diagnosis."""
import glob
import json
import math
import os
import statistics

import analyze_c2_context_validation as context_audit


ROOT = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "results")
PREFIX = "c2_nodepert_val_v1_"
DEPTHS = (1, 4)
TASK_SEEDS = (3, 4, 5)
MODEL_SEEDS = (0, 1, 2, 3, 4)


def audit_nodepert_row(path, row):
    args = row["args"]
    required = {
        "mode": "nodepert", "dataset": "tentmap", "width": 8,
        "act": "relu", "residual": 1, "epochs": 80, "batch_size": 256,
        "eta": 0.03, "momentum": 0.9, "pert_sigma": 0.01,
        "pert_every": 1, "pert_ndirs": 16, "pert_mode": "layerwise",
        "traffic_mode": "none", "nuis_rho": 0.0, "use_residual": 0,
        "learn_A": 0, "learn_P": 0, "p_neutral": 1,
        "task_train_examples": 10000, "task_test_examples": 5000,
        "task_levels": 2, "task_n_in": 1, "val_examples": 2000,
        "split_seed": 2027, "eval_split": "validation", "eval_every": 0,
        "diagnostics": "alignment", "diagnostics_schedule": "final",
        "probe_bs": 512,
    }
    mismatches = {key: (args.get(key), expected) for key, expected in required.items()
                  if args.get(key) != expected}
    expected_lesion = 1.0 / 3.0 if args["depth"] == 4 else 0.0
    if abs(args.get("residual_lesion_fraction", 0.0) - expected_lesion) > 1e-12:
        mismatches["residual_lesion_fraction"] = (
            args.get("residual_lesion_fraction"), expected_lesion)
    if mismatches:
        raise RuntimeError(f"protocol mismatch {path}: {mismatches}")
    if row["final"].get("eval_split") != "validation":
        raise RuntimeError(f"test-contaminated validation row: {path}")
    if any("eval_acc" in step or "cos_q_negg" in step for step in row.get("steps", [])):
        raise RuntimeError(f"intermediate held-out metric/diagnostic: {path}")
    split = row.get("split", {})
    if (not split.get("split_from_training_only")
            or split.get("validation_examples") != 2000
            or split.get("evaluation_split") != "validation"):
        raise RuntimeError(f"invalid validation split {path}: {split}")
    protocol = row.get("diagnostic_protocol", {})
    if protocol != {"probe_source": "training_prefix", "probe_examples": 512,
                    "schedule": "final"}:
        raise RuntimeError(f"diagnostic protocol mismatch {path}: {protocol}")
    if row.get("provenance", {}).get("git_dirty") is not False:
        raise RuntimeError(f"dirty or unknown source provenance: {path}")
    if args["depth"] == 4:
        lesion = row["final"].get("residual_lesion")
        if not lesion or lesion.get("lesioned_layers") != [3]:
            raise RuntimeError(f"incorrect d4 lesion: {path}")


def mean_sd(values):
    return statistics.mean(values), statistics.stdev(values)


def load_context_reference():
    rows = {}
    split_hashes = {}
    for path in sorted(glob.glob(os.path.join(ROOT, context_audit.PREFIX + "*.json"))):
        with open(path) as handle:
            row = json.load(handle)
        context_audit.audit_row(path, row)
        args = row["args"]
        key = (args["mode"], args["depth"], args["task_seed"], args["seed"])
        rows[key] = row
        split_hashes.setdefault(args["task_seed"], set()).add(
            row["split"]["validation_index_sha256"])
    expected = {(method, depth, task_seed, model_seed)
                for method in context_audit.METHODS for depth in DEPTHS
                for task_seed in TASK_SEEDS for model_seed in MODEL_SEEDS}
    if set(rows) != expected:
        raise RuntimeError("context reference panel is incomplete or mixed")
    return rows, split_hashes


def main():
    reference, split_hashes = load_context_reference()
    paths = sorted(glob.glob(os.path.join(ROOT, PREFIX + "*.json")))
    rows = {}
    commits = set()
    for path in paths:
        with open(path) as handle:
            row = json.load(handle)
        audit_nodepert_row(path, row)
        args = row["args"]
        key = (args["depth"], args["task_seed"], args["seed"])
        if key in rows:
            raise RuntimeError(f"duplicate nodepert row: {key}")
        rows[key] = row
        commits.add(row["provenance"]["git_commit"])
        split_hashes.setdefault(args["task_seed"], set()).add(
            row["split"]["validation_index_sha256"])
    expected = {(depth, task_seed, model_seed) for depth in DEPTHS
                for task_seed in TASK_SEEDS for model_seed in MODEL_SEEDS}
    if set(rows) != expected or len(commits) != 1:
        raise RuntimeError(f"incomplete/mixed nodepert panel: rows={len(rows)}, "
                           f"missing={expected - set(rows)}, extra={set(rows) - expected}, "
                           f"commits={commits}")
    if any(len(hashes) != 1 for hashes in split_hashes.values()):
        raise RuntimeError(f"nodepert/reference split mismatch: {split_hashes}")

    accs = {}
    gains = {}
    for method in context_audit.METHODS:
        for depth in DEPTHS:
            accs[(method, depth)] = [
                100 * reference[(method, depth, task_seed, model_seed)]["final"]["val_acc"]
                for task_seed in TASK_SEEDS for model_seed in MODEL_SEEDS]
    for depth in DEPTHS:
        accs[("nodepert", depth)] = [
            100 * rows[(depth, task_seed, model_seed)]["final"]["val_acc"]
            for task_seed in TASK_SEEDS for model_seed in MODEL_SEEDS]
    for method in (*context_audit.METHODS, "nodepert"):
        gains[method] = [deep - shallow for shallow, deep in
                         zip(accs[(method, 1)], accs[(method, 4)])]

    print(f"nodepert_commit={next(iter(commits))} rows={len(rows)} "
          f"task_seeds={list(TASK_SEEDS)}")
    print("| method | depth | validation (%) | depth gain | lesion drop |")
    print("|:---|---:|---:|---:|---:|")
    for method in (*context_audit.METHODS, "nodepert"):
        for depth in DEPTHS:
            mean, sd = mean_sd(accs[(method, depth)])
            gain_text = "--" if depth == 1 else f"{statistics.mean(gains[method]):+.3f}"
            if depth == 1:
                lesion_text = "--"
            else:
                source = rows if method == "nodepert" else reference
                drops = []
                for task_seed in TASK_SEEDS:
                    for model_seed in MODEL_SEEDS:
                        key = ((4, task_seed, model_seed) if method == "nodepert"
                               else (method, 4, task_seed, model_seed))
                        drops.append(100 * source[key]["final"]["residual_lesion"]
                                     ["lesion_acc_drop"])
                lesion_text = f"{statistics.mean(drops):+.3f}"
            print(f"| {method} | {depth} | {mean:.3f} +/- {sd:.3f} | "
                  f"{gain_text} | {lesion_text} |")

    bp_gain = statistics.mean(gains["bp"])
    nodepert_gain = statistics.mean(gains["nodepert"])
    recovery = nodepert_gain / bp_gain if bp_gain > 0 else -math.inf
    positive_gains = sum(value > 0 for value in gains["nodepert"])
    nodepert_d4 = statistics.mean(accs[("nodepert", 4)])
    context_d4 = statistics.mean(accs[("sdil", 4)])
    strongest_fixed = max(statistics.mean(accs[(method, 4)]) for method in ("fa", "dfa"))
    lesion_drops = [100 * rows[(4, task_seed, model_seed)]["final"]
                    ["residual_lesion"]["lesion_acc_drop"]
                    for task_seed in TASK_SEEDS for model_seed in MODEL_SEEDS]
    lesion_mean = statistics.mean(lesion_drops)
    lesion_positive = sum(value > 0 for value in lesion_drops)
    nonfinite = [key for key, row in rows.items()
                 if not math.isfinite(row["final"].get("val_loss", math.nan))]
    d4_rows = [rows[(4, task_seed, model_seed)] for task_seed in TASK_SEEDS
               for model_seed in MODEL_SEEDS]
    forward_work = statistics.mean(
        row["cost"]["training_forward_equivalent_examples"] for row in d4_rows)
    ordinary_work = statistics.mean(
        row["cost"]["ordinary_training_forward_examples"] for row in d4_rows)
    scalar_queries = statistics.mean(
        row["cost"]["calibration_example_loss_evaluations"] for row in d4_rows)
    wall = statistics.mean(row["final"]["wall_s"] for row in d4_rows)

    passed = (not nonfinite and bp_gain >= 5.0 and recovery >= 0.7
              and positive_gains >= 10 and nodepert_d4 - context_d4 >= 2.0
              and lesion_mean >= 2.0 and lesion_positive >= 10)
    print(f"BP gain={bp_gain:+.3f}; direct-NP gain={nodepert_gain:+.3f}; "
          f"recovery={100 * recovery:.1f}%; positive gains={positive_gains}/15")
    print(f"direct-NP d4 minus context-SDIL={nodepert_d4 - context_d4:+.3f}; "
          f"minus strongest FA/DFA={nodepert_d4 - strongest_fixed:+.3f}")
    print(f"direct-NP lesion mean={lesion_mean:+.3f}; positive={lesion_positive}/15; "
          f"nonfinite={nonfinite}")
    print(f"d4 cost: scalar loss queries={scalar_queries:.0f}; "
          f"forward-equivalent examples={forward_work:.0f} "
          f"({forward_work / ordinary_work:.1f}x ordinary); wall={wall:.1f}s")
    print("C2 causal diagnosis: " +
          ("AMORTIZATION BOTTLENECK CONFIRMED" if passed else
           "DIRECT CAUSAL ESTIMATOR/LOCAL OPTIMIZATION ALSO INSUFFICIENT"))
    if not passed:
        raise SystemExit(1)


if __name__ == "__main__":
    main()