summaryrefslogtreecommitdiff
path: root/experiments/analyze_c2_context_validation.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/analyze_c2_context_validation.py')
-rw-r--r--experiments/analyze_c2_context_validation.py151
1 files changed, 151 insertions, 0 deletions
diff --git a/experiments/analyze_c2_context_validation.py b/experiments/analyze_c2_context_validation.py
new file mode 100644
index 0000000..c1ba9b3
--- /dev/null
+++ b/experiments/analyze_c2_context_validation.py
@@ -0,0 +1,151 @@
+"""Audit the independent C2 context-conditioned validation panel."""
+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_context_val_v1_"
+METHODS = ("bp", "fa", "dfa", "sdil")
+DEPTHS = (1, 4)
+TASK_SEEDS = (3, 4, 5)
+MODEL_SEEDS = (0, 1, 2, 3, 4)
+
+
+def audit_row(path, row):
+ args = row["args"]
+ required = {
+ "dataset": "tentmap", "width": 8, "act": "relu", "residual": 1,
+ "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", "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,
+ "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}
+ method = args["mode"]
+ expected_vectorizer = "context_gated" if method == "sdil" else "linear"
+ expected_eta_a = 0.01 if method == "sdil" else 0.02
+ if args.get("vectorizer_mode") != expected_vectorizer:
+ mismatches["vectorizer_mode"] = (args.get("vectorizer_mode"), expected_vectorizer)
+ if args.get("eta_A") != expected_eta_a:
+ mismatches["eta_A"] = (args.get("eta_A"), expected_eta_a)
+ 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 not math.isfinite(row["final"].get("val_loss", math.nan)):
+ raise RuntimeError(f"nonfinite final validation loss: {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 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 = {}
+ for path in paths:
+ with open(path) as handle:
+ row = json.load(handle)
+ audit_row(path, row)
+ args = row["args"]
+ key = (args["mode"], args["depth"], args["task_seed"], args["seed"])
+ if key in rows:
+ raise RuntimeError(f"duplicate 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 = {(method, depth, task_seed, model_seed)
+ for method in METHODS 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 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"methods did not share splits within tasks: {split_hashes}")
+
+ print(f"commit={next(iter(commits))} rows={len(rows)} task_seeds={list(TASK_SEEDS)}")
+ print("| method | depth | validation (%) | depth gain | lesion drop |")
+ print("|:---|---:|---:|---:|---:|")
+ accs = {}
+ gains = {}
+ for method in METHODS:
+ for depth in DEPTHS:
+ accs[(method, depth)] = [
+ 100 * rows[(method, depth, task_seed, model_seed)]["final"]["val_acc"]
+ for task_seed in TASK_SEEDS for model_seed in MODEL_SEEDS]
+ gains[method] = [deep - shallow for shallow, deep in
+ zip(accs[(method, 1)], accs[(method, 4)])]
+ 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:
+ drops = [100 * rows[(method, 4, task_seed, model_seed)]["final"]
+ ["residual_lesion"]["lesion_acc_drop"]
+ for task_seed in TASK_SEEDS for model_seed in MODEL_SEEDS]
+ 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"])
+ sdil_gain = statistics.mean(gains["sdil"])
+ recovery = sdil_gain / bp_gain if bp_gain > 0 else -math.inf
+ competitor_recoveries = {
+ method: statistics.mean(gains[method]) / bp_gain for method in ("fa", "dfa")
+ }
+ strongest_recovery = max(competitor_recoveries.values())
+ strongest_deep = max(statistics.mean(accs[(method, 4)]) for method in ("fa", "dfa"))
+ deep_advantage = statistics.mean(accs[("sdil", 4)]) - strongest_deep
+ positive_gains = sum(value > 0 for value in gains["sdil"])
+ lesion_drops = [100 * rows[("sdil", 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)
+ comparator_ok = strongest_recovery <= 0.5 or deep_advantage >= 2.0
+ passed = (bp_gain >= 5.0 and recovery >= 0.7 and positive_gains >= 10
+ and comparator_ok and lesion_mean >= 2.0 and lesion_positive >= 10)
+
+ print(f"BP gain={bp_gain:+.3f}; context-SDIL gain={sdil_gain:+.3f}; "
+ f"recovery={100 * recovery:.1f}%; positive gains={positive_gains}/15")
+ print(f"competitor recoveries={competitor_recoveries}; strongest={100 * strongest_recovery:.1f}%")
+ print(f"context-SDIL d4 advantage over strongest endpoint={deep_advantage:+.3f}")
+ print(f"context-SDIL lesion mean={lesion_mean:+.3f}; positive={lesion_positive}/15")
+ print(f"C2 context validation gate: {'PASS' if passed else 'FAIL'}")
+ if not passed:
+ raise SystemExit(1)
+
+
+if __name__ == "__main__":
+ main()