summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 05:46:54 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 05:46:54 -0500
commitb4458cdb3533f7b342f3ff3c8e3582962181d09f (patch)
tree213a24676a6d9916c130c0b45cb72a26d41a57b8 /experiments
parent97f49e09d81a9aa332d5d0a07741686e13c65c61 (diff)
experiments: freeze oral B development screen
Diffstat (limited to 'experiments')
-rw-r--r--experiments/analyze_bci_development.py194
-rwxr-xr-xexperiments/bci_development_screen.sh36
2 files changed, 230 insertions, 0 deletions
diff --git a/experiments/analyze_bci_development.py b/experiments/analyze_bci_development.py
new file mode 100644
index 0000000..a7dd259
--- /dev/null
+++ b/experiments/analyze_bci_development.py
@@ -0,0 +1,194 @@
+#!/usr/bin/env python3
+"""Apply the frozen oral-B development eligibility and selection rules."""
+import argparse
+import glob
+import hashlib
+import json
+import math
+import os
+
+
+FEEDBACKS = ("error", "error_velocity")
+KAPPAS = (0.0, 0.1, 0.3)
+ETAS = (0.01, 0.03)
+TASK_SEEDS = (0, 1, 2)
+
+
+def file_hash(path):
+ digest = hashlib.sha256()
+ with open(path, "rb") as handle:
+ for chunk in iter(lambda: handle.read(1024 * 1024), b""):
+ digest.update(chunk)
+ return digest.hexdigest()
+
+
+def require(condition, message):
+ if not condition:
+ raise ValueError(message)
+
+
+def finite(value):
+ return isinstance(value, (int, float)) and math.isfinite(value)
+
+
+def candidate_key(row):
+ args = row["args"]
+ return args["feedback"], float(args["kappa"]), float(args["eta"])
+
+
+def validate_row(row, path):
+ require(row.get("schema_version") == 1, f"{path}: schema mismatch")
+ args = row.get("args", {})
+ require(args.get("task_seed") in TASK_SEEDS, f"{path}: task seed outside development")
+ require(args.get("model_seed") == 0, f"{path}: model seed must be zero")
+ require(args.get("feedback") in FEEDBACKS, f"{path}: feedback mismatch")
+ require(float(args.get("kappa")) in KAPPAS, f"{path}: kappa mismatch")
+ require(float(args.get("eta")) in ETAS, f"{path}: eta mismatch")
+ protocol = row.get("protocol", {})
+ require(protocol.get("name") == "oral_b_continuous_bci_development_v1",
+ f"{path}: protocol mismatch")
+ require(protocol.get("selection_split") == "development",
+ f"{path}: non-development result")
+ require(protocol.get("confirmation_seeds_touched") is False,
+ f"{path}: confirmation leakage")
+ require(protocol.get("evaluation_task_seed") == args["task_seed"] + 100_000,
+ f"{path}: evaluation seed mismatch")
+ provenance = row.get("provenance", {})
+ require(provenance.get("git_dirty") is False, f"{path}: dirty source")
+ require(provenance.get("runner_and_protocol_tracked") is True,
+ f"{path}: untracked runner or protocol")
+ require(isinstance(provenance.get("git_commit"), str)
+ and len(provenance["git_commit"]) == 40,
+ f"{path}: invalid source revision")
+ require(row.get("finite") is True, f"{path}: non-finite run")
+ config = row.get("config", {})
+ expected = {
+ "n_plus": 5, "n_minus": 5, "n_background": 30,
+ "context_dim": 16, "steps_per_episode": 28,
+ "episodes_per_day": 64, "days": 14, "target": 0.8,
+ "perturb_every": 4, "perturb_sigma": 0.03,
+ "feedback": args["feedback"], "kappa": args["kappa"],
+ "forward_eta": args["eta"],
+ }
+ for key, value in expected.items():
+ require(config.get(key) == value,
+ f"{path}: config {key} expected {value!r}, found {config.get(key)!r}")
+ for name in ("intact", "fixed_vectorizer", "online_lesion",
+ "plasticity_lesion", "both_lesion"):
+ condition = row.get("conditions", {}).get(name, {})
+ for key in ("early_success", "late_success", "learning_gain", "final_success"):
+ require(finite(condition.get(key)), f"{path}: {name}.{key} non-finite")
+ for value in row.get("signatures", {}).values():
+ require(finite(value), f"{path}: non-finite signature")
+
+
+def seed_gate(row):
+ intact = row["conditions"]["intact"]
+ fixed = row["conditions"]["fixed_vectorizer"]
+ both = row["conditions"]["both_lesion"]
+ signatures = row["signatures"]
+ checks = {
+ "learning_gain_10pt": intact["learning_gain"] >= 0.10,
+ "fixed_vectorizer_gap_20pt": (
+ intact["final_success"] - fixed["final_success"] >= 0.20),
+ "residual_soma_corr": signatures["mean_abs_residual_soma_corr"] <= 0.10,
+ "positive_sign_inversion": (
+ signatures["causal_role_sign_inversion_index"] > 0),
+ "both_lesion_removes_half_gain": (
+ both["learning_gain"] <= 0.5 * intact["learning_gain"]),
+ }
+ return checks
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--results", default="results/bci_dev")
+ parser.add_argument("--output", default="results/bci_dev_selection.json")
+ args = parser.parse_args()
+ paths = sorted(glob.glob(os.path.join(args.results, "bci_dev_v1_*.json")))
+ require(len(paths) == 36, f"expected 36 development files, found {len(paths)}")
+ groups = {}
+ seen = set()
+ for path in paths:
+ with open(path) as handle:
+ row = json.load(handle)
+ validate_row(row, path)
+ identity = candidate_key(row) + (row["args"]["task_seed"],)
+ require(identity not in seen, f"duplicate development cell {identity}")
+ seen.add(identity)
+ row["_path"] = path
+ groups.setdefault(candidate_key(row), []).append(row)
+ expected = {(feedback, kappa, eta, seed)
+ for feedback in FEEDBACKS for kappa in KAPPAS for eta in ETAS
+ for seed in TASK_SEEDS}
+ require(seen == expected, f"development grid mismatch: missing={sorted(expected-seen)}")
+
+ summaries = []
+ for key in sorted(groups):
+ rows = sorted(groups[key], key=lambda row: row["args"]["task_seed"])
+ checks = [seed_gate(row) for row in rows]
+ eligible = all(all(values.values()) for values in checks)
+ summaries.append({
+ "feedback": key[0], "kappa": key[1], "eta": key[2],
+ "eligible": eligible,
+ "seed_checks": checks,
+ "task_seeds": [row["args"]["task_seed"] for row in rows],
+ "final_success": [row["conditions"]["intact"]["final_success"]
+ for row in rows],
+ "learning_gain": [row["conditions"]["intact"]["learning_gain"]
+ for row in rows],
+ "fixed_final_gap": [
+ row["conditions"]["intact"]["final_success"]
+ - row["conditions"]["fixed_vectorizer"]["final_success"]
+ for row in rows],
+ "sign_inversion": [
+ row["signatures"]["causal_role_sign_inversion_index"] for row in rows],
+ "worst_final_success": min(
+ row["conditions"]["intact"]["final_success"] for row in rows),
+ "worst_sign_inversion": min(
+ row["signatures"]["causal_role_sign_inversion_index"] for row in rows),
+ "source_paths": [row["_path"] for row in rows],
+ "source_commits": sorted({row["provenance"]["git_commit"] for row in rows}),
+ })
+
+ eligible = [summary for summary in summaries if summary["eligible"]]
+ selected = None
+ if eligible:
+ best_final = max(summary["worst_final_success"] for summary in eligible)
+ finalists = [summary for summary in eligible
+ if summary["worst_final_success"] >= best_final - 0.01]
+ finalists.sort(key=lambda summary: (
+ -summary["worst_sign_inversion"], summary["kappa"], summary["eta"],
+ 0 if summary["feedback"] == "error" else 1))
+ chosen = finalists[0]
+ selected = {key: chosen[key] for key in (
+ "feedback", "kappa", "eta", "worst_final_success",
+ "worst_sign_inversion", "source_commits")}
+
+ output = {
+ "schema_version": 1,
+ "protocol": "oral_b_continuous_bci_development_v1",
+ "complete_grid": True,
+ "confirmation_seeds_touched": False,
+ "status": "selected" if selected else "failed_no_eligible_candidate",
+ "selected": selected,
+ "candidates": summaries,
+ "source_sha256": {path: file_hash(path) for path in paths},
+ }
+ if os.path.exists(args.output):
+ raise FileExistsError(f"refusing to overwrite {args.output}")
+ with open(args.output, "w") as handle:
+ json.dump(output, handle, indent=2, sort_keys=True)
+ handle.write("\n")
+ print("feedback kappa eta eligible worst-final worst-sign")
+ for summary in summaries:
+ print(f"{summary['feedback']:<18} {summary['kappa']:>5.1f} "
+ f"{summary['eta']:>5.2f} {str(summary['eligible']):>9} "
+ f"{summary['worst_final_success']:>11.3f} "
+ f"{summary['worst_sign_inversion']:>10.4f}")
+ print(json.dumps({"status": output["status"], "selected": selected},
+ indent=2, sort_keys=True))
+
+
+if __name__ == "__main__":
+ main()
diff --git a/experiments/bci_development_screen.sh b/experiments/bci_development_screen.sh
new file mode 100755
index 0000000..eb4cbba
--- /dev/null
+++ b/experiments/bci_development_screen.sh
@@ -0,0 +1,36 @@
+#!/usr/bin/env bash
+# Frozen 36-run oral-B development screen: 12 candidates x task seeds 0..2.
+set -eu
+
+cd "$(dirname "$0")/.."
+PY=/home/yurenh2/miniconda3/envs/ep_pascal/bin/python3
+OUT=results/bci_dev
+LOG=logs/bci_dev
+mkdir -p "$OUT" "$LOG"
+
+for task_seed in 0 1 2; do
+ for feedback in error error_velocity; do
+ for eta in 0.01 0.03; do
+ for kappa in 0.0 0.1 0.3; do
+ feedback_tag=${feedback/error_velocity/errvel}
+ feedback_tag=${feedback_tag/error/err}
+ eta_tag=${eta/./p}
+ kappa_tag=${kappa/./p}
+ tag="bci_dev_v1_fb${feedback_tag}_k${kappa_tag}_eta${eta_tag}_t${task_seed}_m0"
+ result="$OUT/$tag.json"
+ if [[ -f "$result" ]]; then
+ echo "skip $tag"
+ continue
+ fi
+ echo "run $tag"
+ OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 "$PY" experiments/bci_run.py \
+ --task-seed "$task_seed" --model-seed 0 \
+ --feedback "$feedback" --kappa "$kappa" --eta "$eta" \
+ --tag "$tag" --outdir "$OUT" > "$LOG/$tag.log" 2>&1
+ done
+ done
+ done
+done
+
+"$PY" experiments/analyze_bci_development.py \
+ --results "$OUT" --output results/bci_dev_selection.json