summaryrefslogtreecommitdiff
path: root/experiments/analyze_oral_a_v2_full.py
blob: 16f7e78dc3ea146c8f283802776a603a09bf8830 (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
#!/usr/bin/env python3
"""Apply the frozen Oral-A-v2 full-validation gate."""
import argparse
import json
import math
import os


SPLIT_HASH = "8328b206a97c420e49e54e3eca4abe3274c4756b084355784ea3fb8059e4515b"


def finite_tree(value):
    if isinstance(value, dict):
        return all(finite_tree(item) for item in value.values())
    if isinstance(value, list):
        return all(finite_tree(item) for item in value)
    if isinstance(value, float):
        return math.isfinite(value)
    return True


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument(
        "--run", default="results/oral_a_v2_dev/sdil_full_r20_s0.json")
    parser.add_argument(
        "--selection", default="results/oral_a_v2_calibration_gate.json")
    parser.add_argument(
        "--bp", default="results/oral_a_dev/bp_reference_primary.json")
    parser.add_argument(
        "--dfa", default="results/oral_a_dev/dfa_full_r20_s0.json")
    parser.add_argument("--out", default="results/oral_a_v2_full_gate.json")
    args = parser.parse_args()
    with open(args.run) as handle:
        run = json.load(handle)
    with open(args.selection) as handle:
        selection = json.load(handle)
    with open(args.bp) as handle:
        bp = json.load(handle)
    with open(args.dfa) as handle:
        dfa = json.load(handle)
    if selection["status"] != "passed":
        raise ValueError("v2 causal-capture gate did not pass")
    chosen = selection["selected"]["channel_subspace"]
    expected = {
        "mode": "sdil", "depth": 20, "width": 16, "seed": 0,
        "epochs": 200, "val_examples": 5000, "lr": 0.03,
        "output_lr": 0.1, "lr_schedule": "step",
        "lr_milestones": "100,150", "lr_gamma": 0.1,
        "a_scale": 0.0, "eta_A": chosen["eta_A"],
        "a_warmup_steps": 400, "apical_calibration_mode": "channel_subspace",
        "pert_sigma": 0.01, "pert_directions": 1, "pert_every": 4,
        "normalization": "batchnorm", "vectorizer_mode": "channel_gated",
    }
    for key, value in expected.items():
        if run["args"].get(key) != value:
            raise ValueError(
                f"v2 run {key}={run['args'].get(key)!r}, expected {value!r}")
    if run["provenance"]["git_tracked_dirty"]:
        raise ValueError("tracked-dirty v2 result")
    if run["split"]["validation_index_sha256"] != SPLIT_HASH:
        raise ValueError("v2 split drift")
    if run["evaluation_protocol"]["test_evaluations"] != 0:
        raise ValueError("test endpoint touched during v2 development")
    values = run["diagnostics"]["teaching_negative_gradient_cosine"]
    early_count = max(1, len(values) // 3)
    early = sum(values[:early_count]) / early_count
    sdil_accuracy = run["final"]["accuracy"]
    bp_accuracy = bp["final"]["accuracy"]
    dfa_accuracy = dfa["final"]["accuracy"]
    checks = {
        "all_metrics_finite": run["final"]["finite"] and finite_tree(run),
        "bp_reference_at_least_90pct": bp_accuracy >= 0.90,
        "sdil_within_5pt_of_bp": sdil_accuracy >= bp_accuracy - 0.05,
        "sdil_at_least_2pt_above_dfa": sdil_accuracy >= dfa_accuracy + 0.02,
        "early_third_alignment_at_least_0.05": early >= 0.05,
        "sdil_macs_no_more_than_bp": (
            run["work"]["total_macs_estimate"]
            <= bp["work"]["total_macs_estimate"]),
    }
    passed = all(checks.values())
    output = {
        "protocol": "oral_a_v2_full_validation_v1",
        "status": "passed" if passed else "failed",
        "checks": checks,
        "metrics": {
            "sdil_validation_accuracy": sdil_accuracy,
            "bp_validation_accuracy": bp_accuracy,
            "dfa_validation_accuracy": dfa_accuracy,
            "early_third_alignment": early,
            "sdil_total_macs": run["work"]["total_macs_estimate"],
            "bp_total_macs": bp["work"]["total_macs_estimate"],
        },
        "sources": {"sdil": args.run, "bp": args.bp, "dfa": args.dfa},
        "confirmation_test_seeds_touched": False,
        "review_score_before": 5,
        "review_score_after": 6 if passed else 5,
        "review_score_rationale": (
            "full standard-scale validation gate passed; independent depth "
            "confirmation still required"
            if passed else
            "full standard-scale validation gate failed; controlled evidence "
            "remains unchanged"),
    }
    os.makedirs(os.path.dirname(os.path.abspath(args.out)), exist_ok=True)
    with open(args.out, "w") as handle:
        json.dump(output, handle, indent=2, sort_keys=True)
        handle.write("\n")
    print(json.dumps({
        "status": output["status"], "checks": checks,
        "metrics": output["metrics"],
    }, indent=2))


if __name__ == "__main__":
    main()