diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 05:45:43 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 05:45:43 -0500 |
| commit | 97f49e09d81a9aa332d5d0a07741686e13c65c61 (patch) | |
| tree | b4920b40e718c8929e028310ff2566f7cfdd710b /experiments | |
| parent | c3709dc0dead3234a7423ae705da9703e4d52f68 (diff) | |
bci: add paired phase-lesion runner
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/bci_run.py | 289 |
1 files changed, 289 insertions, 0 deletions
diff --git a/experiments/bci_run.py b/experiments/bci_run.py new file mode 100644 index 0000000..eb7234d --- /dev/null +++ b/experiments/bci_run.py @@ -0,0 +1,289 @@ +#!/usr/bin/env python3 +"""Run one paired continuous-BCI candidate from the frozen oral-B protocol.""" +import argparse +import hashlib +import json +import math +import os +import platform +import resource +import subprocess +import sys +import time + +import torch + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +from sdil.bci import BCIConfig, BCISDIL, generate_trajectories, run_day +from sdil.bci_metrics import annotate_day_events, signature_metrics + + +ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +ORAL_B_PATH = os.path.join(ROOT, "ORAL_B.md") + + +def sha256(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 source_provenance(): + commit = subprocess.run( + ["git", "rev-parse", "HEAD"], cwd=ROOT, check=True, + capture_output=True, text=True).stdout.strip() + dirty = bool(subprocess.run( + ["git", "status", "--porcelain", "--untracked-files=no"], cwd=ROOT, + check=True, capture_output=True, text=True).stdout.strip()) + runner_path = os.path.relpath(os.path.abspath(__file__), ROOT) + protocol_path = os.path.relpath(ORAL_B_PATH, ROOT) + tracked = all(subprocess.run( + ["git", "ls-files", "--error-unmatch", path], cwd=ROOT, + capture_output=True).returncode == 0 for path in (runner_path, protocol_path)) + return {"git_commit": commit, "git_dirty": dirty, + "runner_and_protocol_tracked": tracked, + "oral_b_protocol_sha256": sha256(ORAL_B_PATH)} + + +def finite_tree(value): + if isinstance(value, dict): + return all(finite_tree(item) for item in value.values()) + if isinstance(value, (list, tuple)): + return all(finite_tree(item) for item in value) + if isinstance(value, (int, float)): + return math.isfinite(value) + return True + + +def condition_settings(name): + settings = { + "intact": dict(control_gain=1.0, plasticity_gain=1.0, + learn_vectorizer=True), + "fixed_vectorizer": dict(control_gain=1.0, plasticity_gain=1.0, + learn_vectorizer=False), + "online_lesion": dict(control_gain=0.0, plasticity_gain=1.0, + learn_vectorizer=True), + "plasticity_lesion": dict(control_gain=1.0, plasticity_gain=0.0, + learn_vectorizer=True), + "both_lesion": dict(control_gain=0.0, plasticity_gain=0.0, + learn_vectorizer=True), + } + return settings[name] + + +def train_condition(initial, trajectories, name, *, collect): + model = initial.clone() + settings = condition_settings(name) + daily_success = [] + events = [] + global_step = 0 + perturbation_events = 0 + perturbation_scalar_rewards = 0 + start = time.perf_counter() + for day in range(model.cfg.days): + report = run_day( + model, trajectories, day, collect=collect, global_step=global_step, + control_gain=settings["control_gain"], + plasticity_gain=settings["plasticity_gain"], + learn_vectorizer=settings["learn_vectorizer"], + learn_predictor=True) + global_step = report["global_step"] + daily_success.append(report["success_rate"]) + if collect: + annotate_day_events( + report["events"], day, report["success"], + episode_offset=day * model.cfg.episodes_per_day) + events.extend(report["events"]) + for event in report["events"]: + if event.get("causal_target") is not None: + perturbation_events += 1 + perturbation_scalar_rewards += 2 * int(event["active"].sum()) + elif settings["learn_vectorizer"]: + # One batched event every perturb_every temporal steps. + perturbation_events += sum( + 1 for step in range(model.cfg.steps_per_episode) + if ((day * model.cfg.steps_per_episode + step) + % model.cfg.perturb_every == 0)) + # The exact active count is unavailable without collection. Report + # the conservative full-batch count for non-primary conditions. + perturbation_scalar_rewards += ( + 2 * model.cfg.episodes_per_day + * sum(1 for step in range(model.cfg.steps_per_episode) + if ((day * model.cfg.steps_per_episode + step) + % model.cfg.perturb_every == 0))) + elapsed = time.perf_counter() - start + early = sum(daily_success[:3]) / 3 + late = sum(daily_success[-3:]) / 3 + ordinary = (model.cfg.days * model.cfg.episodes_per_day + * model.cfg.steps_per_episode) + return model, events, { + "daily_success": daily_success, + "early_success": early, + "late_success": late, + "learning_gain": late - early, + "training_wall_s": elapsed, + "cost": { + "ordinary_state_episode_steps": ordinary, + "perturbation_events": perturbation_events, + "causal_scalar_reward_observations": perturbation_scalar_rewards, + "forward_equivalent_episode_steps": ( + ordinary + perturbation_scalar_rewards / model.cfg.context_dim), + "forward_equivalent_definition": ( + "ordinary full state step plus each perturbed cursor/loss readout " + "at 1/context_dim of a context-to-soma state step"), + }, + } + + +def evaluate(model, trajectories, control_gain, *, collect): + # Clone prevents even accidental evaluator mutation from changing another + # paired phase condition. + evaluated = model.clone() + start = time.perf_counter() + report = run_day( + evaluated, trajectories, 0, collect=collect, global_step=0, + control_gain=control_gain, plasticity_gain=0.0, + learn_vectorizer=False, learn_predictor=False) + wall = time.perf_counter() - start + if collect: + annotate_day_events(report["events"], 0, report["success"], 1_000_000) + return report, wall + + +def build_config(args): + return BCIConfig( + 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, inertia=0.65, process_noise=0.12, context_ar=0.8, + coupling_scale=1.0, predictor_eta=0.2, vectorizer_eta=0.03, + forward_eta=args.eta, perturb_sigma=0.03, perturb_every=4, + kappa=args.kappa, feedback=args.feedback) + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--task-seed", type=int, required=True) + parser.add_argument("--model-seed", type=int, required=True) + parser.add_argument("--feedback", choices=("error", "error_velocity"), required=True) + parser.add_argument("--kappa", type=float, choices=(0.0, 0.1, 0.3), required=True) + parser.add_argument("--eta", type=float, choices=(0.01, 0.03), required=True) + parser.add_argument("--tag", required=True) + parser.add_argument("--outdir", default="results/bci_dev") + args = parser.parse_args() + if args.task_seed not in (0, 1, 2): + raise ValueError("development runner permits only frozen task seeds 0,1,2") + if args.model_seed != 0: + raise ValueError("development runner permits only frozen model seed 0") + + provenance = source_provenance() + if provenance["git_dirty"] or not provenance["runner_and_protocol_tracked"]: + raise RuntimeError( + "refusing to generate oral-B evidence from dirty or untracked source") + cfg = build_config(args) + train_trajectories = generate_trajectories(cfg, args.task_seed) + evaluation_trajectories = generate_trajectories( + cfg, args.task_seed + 100_000, days=1, episodes=256) + initial = BCISDIL(cfg, args.model_seed) + started = time.perf_counter() + + trained = {} + training_events = None + condition_reports = {} + for name in ("intact", "fixed_vectorizer", "online_lesion", + "plasticity_lesion", "both_lesion"): + model, events, report = train_condition( + initial, train_trajectories, name, collect=name == "intact") + trained[name] = model + condition_reports[name] = report + if name == "intact": + training_events = events + + evaluation_events = None + for name, model in trained.items(): + control_gain = condition_settings(name)["control_gain"] + report, wall = evaluate( + model, evaluation_trajectories, control_gain, + collect=name == "intact") + condition_reports[name]["final_success"] = report["success_rate"] + condition_reports[name]["evaluation_wall_s"] = wall + if name == "intact": + evaluation_events = report["events"] + + acute, acute_wall = evaluate( + trained["intact"], evaluation_trajectories, 0.0, collect=False) + sham, sham_wall = evaluate( + trained["intact"], evaluation_trajectories, 1.0, collect=False) + intact_final = condition_reports["intact"]["final_success"] + if sham["success_rate"] != intact_final: + raise RuntimeError("paired sham evaluation changed intact success") + + signatures = signature_metrics( + training_events, evaluation_events, cfg, trained["intact"].role) + condition_reports["acute_online_lesion_after_intact"] = { + "final_success": acute["success_rate"], + "drop_from_intact": intact_final - acute["success_rate"], + "evaluation_wall_s": acute_wall, + } + condition_reports["sham_after_intact"] = { + "final_success": sham["success_rate"], + "change_from_intact": sham["success_rate"] - intact_final, + "evaluation_wall_s": sham_wall, + } + condition_reports["intact"]["gain_over_fixed_vectorizer_final"] = ( + intact_final - condition_reports["fixed_vectorizer"]["final_success"]) + condition_reports["both_lesion"]["fraction_of_intact_learning_gain"] = ( + condition_reports["both_lesion"]["learning_gain"] + / max(1e-12, condition_reports["intact"]["learning_gain"])) + + result = { + "schema_version": 1, + "args": vars(args), + "config": vars(cfg), + "protocol": { + "name": "oral_b_continuous_bci_development_v1", + "selection_split": "development", + "training_task_seed": args.task_seed, + "evaluation_task_seed": args.task_seed + 100_000, + "confirmation_seeds_touched": False, + "paired_environment_across_conditions": True, + "evaluation_plasticity": False, + }, + "provenance": provenance, + "conditions": condition_reports, + "signatures": signatures, + "finite": None, + "wall_s": time.perf_counter() - started, + "peak_rss_mib": resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / 1024.0, + "hardware": { + "device": "cpu", + "platform": platform.platform(), + "processor": platform.processor(), + "torch_version": torch.__version__, + "threads": torch.get_num_threads(), + }, + } + result["finite"] = finite_tree(result) + if not result["finite"]: + raise RuntimeError("non-finite continuous-BCI result") + os.makedirs(args.outdir, exist_ok=True) + path = os.path.join(args.outdir, args.tag + ".json") + if os.path.exists(path): + raise FileExistsError(f"refusing to overwrite {path}") + with open(path, "w") as handle: + json.dump(result, handle, indent=2, sort_keys=True) + handle.write("\n") + print(json.dumps({ + "path": path, + "intact_final": intact_final, + "intact_gain": condition_reports["intact"]["learning_gain"], + "fixed_final": condition_reports["fixed_vectorizer"]["final_success"], + "both_gain_fraction": condition_reports["both_lesion"][ + "fraction_of_intact_learning_gain"], + "signatures": signatures, + }, indent=2, sort_keys=True)) + + +if __name__ == "__main__": + main() |
