summaryrefslogtreecommitdiff
path: root/experiments/analyze_c2_nlms_pilot.py
blob: f7ed39ac44c8ad152b7e8e1bfbd4fc5d0232385c (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
"""Audit and select the frozen normalized-regression C2 development screen."""
import glob
import json
import math
import os
import statistics


ROOT = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "results")
PREFIX = "c2_nlmsdev_v1_"
ETA_AS = (0.01, 0.05)
DEPTHS = (1, 4)
MODEL_SEEDS = (0, 1, 2)


def audit_row(path, row):
    args = row["args"]
    required = {
        "mode": "sdil", "dataset": "tentmap", "width": 8,
        "act": "relu", "residual": 1, "vectorizer_mode": "context_gated",
        "vectorizer_optimizer": "nlms", "vectorizer_eps": 1e-6,
        "epochs": 80, "batch_size": 256, "eta": 0.03, "momentum": 0.9,
        "eta_P": 0.002, "pert_sigma": 0.01, "pert_every": 4,
        "pert_ndirs": 1, "pert_mode": "simultaneous",
        "a_warmup_steps": 100, "a_warmup_ndirs": 4,
        "a_warmup_mode": "layerwise", "traffic_mode": "none",
        "nuis_rho": 0.0, "use_residual": 1, "learn_A": 1, "learn_P": 1,
        "p_neutral": 1, "task_train_examples": 10000,
        "task_test_examples": 5000, "task_levels": 2, "task_n_in": 1,
        "task_seed": 0, "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}
    if args.get("eta_A") not in ETA_AS:
        mismatches["eta_A"] = (args.get("eta_A"), ETA_AS)
    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 development row: {path}")
    if any("eval_acc" in step or "cos_r_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}")
    if row.get("provenance", {}).get("git_dirty") is not False:
        raise RuntimeError(f"dirty or unknown source provenance: {path}")
    if row.get("cost", {}).get("feedback_warmup_perturbation_events") != 100:
        raise RuntimeError(f"warmup cost mismatch {path}: {row.get('cost')}")
    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 main():
    paths = sorted(glob.glob(os.path.join(ROOT, PREFIX + "*.json")))
    rows = {}
    commits = set()
    split_hashes = set()
    for path in paths:
        with open(path) as handle:
            row = json.load(handle)
        audit_row(path, row)
        args = row["args"]
        key = (args["eta_A"], args["depth"], args["seed"])
        if key in rows:
            raise RuntimeError(f"duplicate row: {key}")
        rows[key] = row
        commits.add(row["provenance"]["git_commit"])
        split_hashes.add(row["split"]["validation_index_sha256"])
    expected = {(eta_a, depth, seed) for eta_a in ETA_AS
                for depth in DEPTHS for seed in MODEL_SEEDS}
    if set(rows) != expected or len(commits) != 1 or len(split_hashes) != 1:
        raise RuntimeError(f"incomplete/mixed pilot: rows={len(rows)}, "
                           f"missing={expected - set(rows)}, extra={set(rows) - expected}, "
                           f"commits={commits}, split_hashes={split_hashes}")

    summaries = {}
    print(f"commit={next(iter(commits))} rows={len(rows)} task_seed=0")
    print("| eta_A | depth | validation (%) | depth gain | d4 lesion | d4 work/ordinary |")
    print("|---:|---:|---:|---:|---:|---:|")
    for eta_a in ETA_AS:
        accs = {depth: [100 * rows[(eta_a, depth, seed)]["final"]["val_acc"]
                        for seed in MODEL_SEEDS] for depth in DEPTHS}
        gains = [deep - shallow for shallow, deep in zip(accs[1], accs[4])]
        nonfinite = [seed for depth in DEPTHS for seed in MODEL_SEEDS
                     if not math.isfinite(rows[(eta_a, depth, seed)]["final"]
                                          .get("val_loss", math.nan))]
        lesions = [100 * rows[(eta_a, 4, seed)]["final"]["residual_lesion"]
                   ["lesion_acc_drop"] for seed in MODEL_SEEDS]
        work_ratios = [rows[(eta_a, 4, seed)]["cost"]
                       ["training_forward_equivalent_examples"]
                       / rows[(eta_a, 4, seed)]["cost"]
                       ["ordinary_training_forward_examples"] for seed in MODEL_SEEDS]
        eligible = (not nonfinite and statistics.mean(accs[4]) >= 90.0
                    and sum(value > 0 for value in gains) >= 2
                    and statistics.mean(lesions) >= 2.0
                    and sum(value > 0 for value in lesions) >= 2)
        summaries[eta_a] = {"d4_mean": statistics.mean(accs[4]), "eligible": eligible}
        for depth in DEPTHS:
            mean, sd = mean_sd(accs[depth])
            gain_text = "--" if depth == 1 else f"{statistics.mean(gains):+.3f}"
            lesion_text = "--" if depth == 1 else f"{statistics.mean(lesions):+.3f}"
            work_text = "--" if depth == 1 else f"{statistics.mean(work_ratios):.1f}x"
            print(f"| {eta_a:.2f} | {depth} | {mean:.3f} +/- {sd:.3f} | "
                  f"{gain_text} | {lesion_text} | {work_text} |")
        print(f"eta_A={eta_a:.2f} positive gains={sum(value > 0 for value in gains)}/3; "
              f"positive lesions={sum(value > 0 for value in lesions)}/3; "
              f"nonfinite={nonfinite}; eligible={eligible}")

    eligible = [eta_a for eta_a in ETA_AS if summaries[eta_a]["eligible"]]
    if not eligible:
        print("C2 normalized-regression development: NO ELIGIBLE CANDIDATE")
        raise SystemExit(1)
    best_mean = max(summaries[eta_a]["d4_mean"] for eta_a in eligible)
    selected = min(eta_a for eta_a in eligible
                   if summaries[eta_a]["d4_mean"] >= best_mean - 1.0)
    print(f"C2 normalized-regression development selection: eta_A={selected:.2f} "
          f"(best d4={best_mean:.3f}%, selected d4={summaries[selected]['d4_mean']:.3f}%)")


if __name__ == "__main__":
    main()