diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-06 14:19:26 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-06 14:19:26 -0500 |
| commit | b3457848820d0818840e7052c41f01c193d04a67 (patch) | |
| tree | c0ad78c897787035cea3982f71ccf5b0855ecb9e /experiments/shared_feedback_s0.py | |
| parent | 20a568d28375a8477eb7f55c533ac2338756ba59 (diff) | |
experiment: implement shared-feedback feasibility gate
Diffstat (limited to 'experiments/shared_feedback_s0.py')
| -rw-r--r-- | experiments/shared_feedback_s0.py | 162 |
1 files changed, 162 insertions, 0 deletions
diff --git a/experiments/shared_feedback_s0.py b/experiments/shared_feedback_s0.py new file mode 100644 index 0000000..719a333 --- /dev/null +++ b/experiments/shared_feedback_s0.py @@ -0,0 +1,162 @@ +#!/usr/bin/env python3 +"""Frozen single-run S0 screen from SHARED_FEEDBACK.md.""" + +import argparse +import json +import os +from pathlib import Path +import subprocess +import sys +import time + +import torch + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +from sdil.shared_feedback import ( + CONDITIONS, SharedFeedbackConfig, SharedFeedbackNet, + conditional_selector_data, evaluate_shared_feedback, shared_feedback_step, +) + + +ROOT = Path(__file__).resolve().parents[1] +DEFAULT_OUT = ROOT / "results" / "shared_feedback" / "s0.json" + + +def git_output(*args): + return subprocess.run( + ["git", *args], cwd=ROOT, check=True, capture_output=True, + text=True).stdout.strip() + + +def train_condition(condition, base, train, validation, neutral, device): + net = base.clone() + x_train, z_train, y_train = train + x_val, z_val, y_val = validation + x_neutral, z_neutral = neutral + shuffle = torch.Generator(device="cpu").manual_seed(3101) + epoch_losses = [] + predictor_reports = [] + first_nonfinite_epoch = None + started = time.time() + for epoch in range(40): + if condition in ("innovation", "matched_raw"): + predictor_reports = net.fit_neutral_predictor(x_neutral, z_neutral) + permutation = torch.randperm(x_train.shape[0], generator=shuffle) + losses = [] + for start in range(0, x_train.shape[0], 128): + indices = permutation[start:start + 128].to(device) + loss, _ = shared_feedback_step( + net, x_train[indices], z_train[indices], y_train[indices], condition) + losses.append(loss) + mean_loss = sum(losses) / len(losses) + epoch_losses.append(mean_loss) + if not torch.isfinite(torch.tensor(mean_loss)): + first_nonfinite_epoch = epoch + break + if condition not in ("innovation", "matched_raw"): + predictor_reports = net.fit_neutral_predictor(x_neutral, z_neutral) + endpoint = evaluate_shared_feedback(net, x_val, z_val, y_val) + lesion = evaluate_shared_feedback( + net, x_val, z_val, y_val, context_enabled=False) + return { + "condition": condition, + "epochs_completed": len(epoch_losses), + "epoch_train_loss": epoch_losses, + "first_nonfinite_epoch": first_nonfinite_epoch, + "finite": first_nonfinite_epoch is None, + "validation": endpoint, + "context_lesion_validation": lesion, + "predictor": predictor_reports, + "wall_seconds": time.time() - started, + } + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--device", default="cpu") + parser.add_argument("--out", type=Path, default=DEFAULT_OUT) + args = parser.parse_args() + if git_output("status", "--porcelain", "--untracked-files=no"): + raise RuntimeError("S0 requires clean tracked source") + device = torch.device(args.device) + torch.manual_seed(3101) + if device.type == "cpu": + torch.set_num_threads(1) + + config = SharedFeedbackConfig() + base = SharedFeedbackNet(config, seed=3101, device=device) + train = conditional_selector_data(8192, 3101, device) + validation = conditional_selector_data(2048, 3102, device) + neutral = (train[0][:512], train[1][:512]) + records = [train_condition( + condition, base, train, validation, neutral, device) + for condition in CONDITIONS] + by_name = {row["condition"]: row for row in records} + oracle = 100.0 * by_name["oracle"]["validation"]["accuracy"] + raw = 100.0 * by_name["raw_shared"]["validation"]["accuracy"] + innovation = 100.0 * by_name["innovation"]["validation"]["accuracy"] + matched = 100.0 * by_name["matched_raw"]["validation"]["accuracy"] + lesion_drop = 100.0 * ( + by_name["oracle"]["validation"]["accuracy"] + - by_name["oracle"]["context_lesion_validation"]["accuracy"]) + predictor = by_name["innovation"]["predictor"] + checks = { + "oracle_at_least_90": oracle >= 90.0, + "oracle_context_lesion_drop_at_least_10": lesion_drop >= 10.0, + "nonzero_context_every_layer": all( + row["context_rms"] > 0 for row in predictor), + "raw_below_oracle_by_5_or_nonfinite": ( + oracle - raw >= 5.0 or not by_name["raw_shared"]["finite"]), + "innovation_above_raw_by_5": innovation - raw >= 5.0, + "innovation_within_3_of_oracle": oracle - innovation <= 3.0, + "innovation_within_2_of_exact_subtraction": abs( + innovation - oracle) <= 2.0, + "matched_raw_below_innovation_by_3": innovation - matched >= 3.0, + "predictor_mean_r2_at_least_0p8": ( + sum(row["mean_per_cell_r2"] for row in predictor) + / len(predictor) >= 0.8), + "predictor_residual_ratio_at_most_0p25": max( + row["residual_context_rms_ratio"] for row in predictor) <= 0.25, + "zero_instruction_observations": max( + row["instruction_observations"] for record in records + for row in record["predictor"]) == 0, + } + report = { + "stage": "shared_feedback_s0", + "gate": "pass" if all(checks.values()) else "fail", + "checks": checks, + "config": config.__dict__, + "data": { + "train_examples": 8192, "validation_examples": 2048, + "neutral_examples_per_epoch": 512, "batch_size": 128, + "epochs": 40, "data_seed": 3101, + "validation_seed": 3102, "test_generated": False, + }, + "records": records, + "summary": { + "oracle_validation_accuracy_percent": oracle, + "raw_validation_accuracy_percent": raw, + "innovation_validation_accuracy_percent": innovation, + "matched_raw_validation_accuracy_percent": matched, + "oracle_context_lesion_drop_points": lesion_drop, + }, + "provenance": { + "git_commit": git_output("rev-parse", "HEAD"), + "git_dirty_tracked": False, + "device": str(device), "torch_version": torch.__version__, + "cuda_device_name": ( + torch.cuda.get_device_name(device) if device.type == "cuda" else None), + "cuda_visible_devices": os.environ.get("CUDA_VISIBLE_DEVICES"), + }, + } + args.out.parent.mkdir(parents=True, exist_ok=True) + with open(args.out, "w", encoding="utf-8") as handle: + json.dump(report, handle, indent=2, sort_keys=True) + handle.write("\n") + print(json.dumps({"gate": report["gate"], **report["summary"], + "checks": checks}, indent=2, sort_keys=True)) + + +if __name__ == "__main__": + main() + |
