From b3457848820d0818840e7052c41f01c193d04a67 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Thu, 6 Aug 2026 14:19:26 -0500 Subject: experiment: implement shared-feedback feasibility gate --- experiments/shared_feedback_s0.py | 162 +++++++++++++++++++++++++++++++++++ experiments/shared_feedback_smoke.py | 103 ++++++++++++++++++++++ 2 files changed, 265 insertions(+) create mode 100644 experiments/shared_feedback_s0.py create mode 100644 experiments/shared_feedback_smoke.py (limited to 'experiments') 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() + diff --git a/experiments/shared_feedback_smoke.py b/experiments/shared_feedback_smoke.py new file mode 100644 index 0000000..672d123 --- /dev/null +++ b/experiments/shared_feedback_smoke.py @@ -0,0 +1,103 @@ +#!/usr/bin/env python3 +"""Deterministic mechanics checks for SHARED_FEEDBACK.md S0.""" + +import os +import sys + +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, select_shared_signal, shared_feedback_step, +) + + +def maximum_difference(left, right): + return max(float((a - b).abs().max()) for a, b in zip(left, right)) + + +def main(): + torch.set_num_threads(1) + config = SharedFeedbackConfig(width=8, hidden_layers=2) + base = SharedFeedbackNet(config, seed=3101, dtype=torch.float64) + x, z, y = conditional_selector_data(32, 3101) + x = x.to(torch.float64) + + clones = {condition: base.clone() for condition in CONDITIONS} + forward_error = max(float(( + clones["oracle"].forward(x, z)["logits"] + - clones[condition].forward(x, z)["logits"]).abs().max()) + for condition in CONDITIONS) + assert forward_error == 0.0 + + enabled = base.forward(x, z) + lesioned = base.forward(x, z, context_enabled=False) + lesion_context_error = max(float(value.abs().max()) + for value in lesioned["context"]) + assert lesion_context_error == 0.0 + assert any(not torch.equal(a, b) for a, b in zip( + enabled["h"][1:-1], lesioned["h"][1:-1])) + + reports = base.fit_neutral_predictor(x, z) + assert all(row["instruction_observations"] == 0 for row in reports) + layer = 0 + soma = base.forward(x, z)["h"][1] + context = base.context_fields(z)[layer] + instruction = torch.randn_like(context) + exact, raw, innovation = select_shared_signal( + base, layer, instruction, context, soma, "exact_subtraction") + assert torch.equal(raw, instruction + context) + assert torch.allclose(exact, instruction, atol=1e-15, rtol=1e-15) + matched, _, _ = select_shared_signal( + base, layer, instruction, context, soma, "matched_raw") + assert torch.allclose(matched.norm(dim=1), innovation.norm(dim=1), + atol=1e-12, rtol=1e-12) + + # With context and predictor exactly zero, all four signal rules coincide. + zero = base.clone() + for value in zero.C + zero.P + zero.P_bias: + value.zero_() + zero_clones = {condition: zero.clone() for condition in CONDITIONS} + for condition, net in zero_clones.items(): + shared_feedback_step(net, x, z, y, condition) + reference = zero_clones["oracle"] + zero_context_update_error = max( + maximum_difference(reference.W, net.W) + + maximum_difference(reference.Q[1:], net.Q[1:]) + for condition, net in zero_clones.items() if condition != "oracle") + assert zero_context_update_error < 1e-14 + + # The reciprocal change is independently applied from the same local + # direction; it is not copied from the updated forward tensor. + kp = base.clone() + before_w = [value.clone() for value in kp.W] + before_q = [None] + [value.clone() for value in kp.Q[1:]] + _, auxiliary = shared_feedback_step(kp, x, z, y, "oracle") + reciprocal_direction_error = 0.0 + for layer in range(1, len(kp.W)): + expected_w = ((kp.W[layer] - before_w[layer]) + / kp.config.learning_rate) + expected_q = ((kp.Q[layer] - before_q[layer]) + / kp.config.reciprocal_learning_rate) + reciprocal_direction_error = max( + reciprocal_direction_error, + float((expected_w - expected_q + - kp.config.weight_decay + * (before_q[layer] - before_w[layer])).abs().max())) + assert reciprocal_direction_error < 1e-12 + assert all(not value.requires_grad for value in kp.W + kp.Q[1:]) + assert all(direction.grad_fn is None for direction in auxiliary["directions"]) + + print({ + "forward_identity_error": forward_error, + "lesion_context_error": lesion_context_error, + "zero_context_update_error": zero_context_update_error, + "reciprocal_direction_error": reciprocal_direction_error, + "neutral_instruction_observations": max( + row["instruction_observations"] for row in reports), + }) + + +if __name__ == "__main__": + main() -- cgit v1.2.3