summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 17:32:54 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 17:32:54 -0500
commitf8fb86a5516997f1b720f9fd24921c5a4a8f1036 (patch)
treeae7ab6a6f36a43b2c10789e9898b5e588cb65929 /experiments
parent2677f429c224f3076af063ecdfa65695eb76c531 (diff)
figure: visualize dynamic innovation stability repair
Diffstat (limited to 'experiments')
-rwxr-xr-xexperiments/plot_dynamic_stability.py220
1 files changed, 220 insertions, 0 deletions
diff --git a/experiments/plot_dynamic_stability.py b/experiments/plot_dynamic_stability.py
new file mode 100755
index 0000000..3c99260
--- /dev/null
+++ b/experiments/plot_dynamic_stability.py
@@ -0,0 +1,220 @@
+#!/usr/bin/env python3
+"""Render the audited dynamic neutral-projection stability figure."""
+import argparse
+import hashlib
+import json
+import math
+import os
+
+import matplotlib
+matplotlib.use("Agg")
+import matplotlib.pyplot as plt
+
+
+PDF_METADATA = {
+ "Creator": "SDIL audited figure pipeline",
+ "Producer": "Matplotlib",
+ "CreationDate": None,
+ "ModDate": None,
+}
+
+
+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 read(path):
+ with open(path) as handle:
+ return json.load(handle)
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument(
+ "--dynamic_train", default="results/kp_dynamic_projection/dynamic.json")
+ parser.add_argument(
+ "--dynamic_train_gate", default="results/kp_dynamic_projection_gate.json")
+ parser.add_argument(
+ "--fixed_lesion",
+ default=("results/kp_innovation_stability/"
+ "closed_form_frozen_innovation.json"))
+ parser.add_argument(
+ "--dynamic_short",
+ default="results/kp_dynamic_projection_short/dynamic.json")
+ parser.add_argument(
+ "--dynamic_short_gate",
+ default="results/kp_dynamic_projection_short_gate.json")
+ parser.add_argument(
+ "--mt1_gate", default="results/kp_innovation_short_gate.json")
+ parser.add_argument(
+ "--kp_gate", default="results/kp_short_gate.json")
+ parser.add_argument(
+ "--out", default="results/figs/figureS_dynamic_stability")
+ args = parser.parse_args()
+
+ dynamic = read(args.dynamic_train)
+ dynamic_gate = read(args.dynamic_train_gate)
+ lesion = read(args.fixed_lesion)
+ short = read(args.dynamic_short)
+ short_gate = read(args.dynamic_short_gate)
+ mt1 = read(args.mt1_gate)
+ kp = read(args.kp_gate)
+ if (dynamic_gate.get("protocol") !=
+ "kp_dynamic_neutral_projection_training_prefix_v1"
+ or dynamic_gate.get("status") != "passed"):
+ raise ValueError("dynamic training gate did not pass")
+ if (short_gate.get("protocol") !=
+ "kp_dynamic_neutral_projection_short_v1"
+ or short_gate.get("status") != "passed"):
+ raise ValueError("dynamic short gate did not pass")
+ if (mt1.get("protocol") != "kp_mixed_traffic_short_v1"
+ or mt1.get("status") != "failed"):
+ raise ValueError("MT-1 comparator is not the frozen failure")
+ if kp.get("protocol") != "kolen_pollack_short_v1" or kp.get(
+ "status") != "passed":
+ raise ValueError("clean KP comparator is not frozen")
+ if dynamic_gate["source_commit"] != dynamic["provenance"]["git_commit"]:
+ raise ValueError("D1 record/gate provenance mismatch")
+ if short_gate["metrics"]["source_commit"] != short[
+ "provenance"]["git_commit"]:
+ raise ValueError("D2 record/gate provenance mismatch")
+ if (lesion["predictor_mode"] != "closed_form"
+ or lesion["predictor_every"] != 0
+ or lesion.get("stability_margin", 0.0) != 0.0
+ or lesion["provenance"]["git_tracked_dirty"]):
+ raise ValueError("fixed-predictor diagnostic drift")
+ if (len(dynamic["trajectory"]) != 352
+ or len(lesion["trajectory"]) != 352):
+ raise ValueError("training-prefix trajectory drift")
+
+ steps = list(range(1, 353))
+ dynamic_loss = [float(row["batch_loss"])
+ for row in dynamic["trajectory"]]
+ lesion_loss = [float(row["batch_loss"])
+ for row in lesion["trajectory"]]
+ pre_ratio = [float(row["neutral_projection"][
+ "pre_projection_traffic_rms_ratio"])
+ for row in dynamic["trajectory"]]
+ post_ratio = [float(row["neutral_projection"][
+ "post_projection_traffic_rms_ratio"])
+ for row in dynamic["trajectory"]]
+ numbers = dynamic_loss + lesion_loss + pre_ratio + post_ratio
+ if not all(math.isfinite(value) and value >= 0 for value in numbers):
+ raise ValueError("nonfinite plotting trajectory")
+
+ accuracies = [
+ 100 * float(mt1["metrics"]["accuracy"]["raw"]),
+ 100 * float(mt1["metrics"]["accuracy"]["matched"]),
+ 100 * float(mt1["metrics"]["accuracy"]["innovation"]),
+ 100 * float(kp["metrics"]["accuracy"]),
+ 100 * float(short_gate["metrics"]["accuracy"]),
+ ]
+ labels = ["Raw", "Norm-\nmatched", "Frozen\ninnovation",
+ "Clean\nKP", "Dynamic\ninnovation"]
+
+ plt.rcParams.update({
+ "font.family": "DejaVu Sans",
+ "font.size": 10,
+ "axes.titleweight": "bold",
+ "axes.spines.top": False,
+ "axes.spines.right": False,
+ })
+ blue = "#0072B2"
+ red = "#C44E52"
+ orange = "#E69F00"
+ gray = "#777777"
+ fig, axes = plt.subplots(1, 3, figsize=(13.2, 4.15),
+ gridspec_kw={"width_ratios": [1.2, 1.05, 1.05]})
+
+ left, middle, right = axes
+ left.semilogy(steps, lesion_loss, color=red, linewidth=1.8,
+ label="Frozen predictor (controller lesion)")
+ left.semilogy(steps, dynamic_loss, color=blue, linewidth=2.1,
+ label="Dynamic neutral projection")
+ left.axhline(10, color="#BBBBBB", linestyle="--", linewidth=0.9)
+ left.set_xlim(1, 352)
+ left.set_xlabel("Training minibatch")
+ left.set_ylabel("Cross-entropy (log scale)")
+ left.set_title("a Stability is a trajectory property", loc="left")
+ left.grid(axis="y", which="major", color="#E1E1E1", linewidth=0.7)
+ left.legend(frameon=False, fontsize=8.5, loc="upper left")
+ left.text(345, lesion_loss[-1], f" {lesion_loss[-1]:.1e}",
+ color=red, fontsize=8.5, va="center")
+
+ middle.semilogy(steps, pre_ratio, color=orange, linewidth=1.9,
+ label="Before fast projection")
+ middle.semilogy(steps, post_ratio, color=blue, linewidth=2.1,
+ label="After fast projection")
+ middle.fill_between(steps, post_ratio, pre_ratio, color=blue, alpha=0.08)
+ middle.set_xlim(1, 352)
+ middle.set_ylim(1e-9, 2e-2)
+ middle.set_xlabel("Training minibatch")
+ middle.set_ylabel("Neutral residual / traffic RMS")
+ middle.set_title("b Local coupling is continually removed", loc="left")
+ middle.grid(axis="y", which="major", color="#E1E1E1", linewidth=0.7)
+ middle.legend(frameon=False, fontsize=8.5, loc="upper left")
+ middle.annotate(
+ f"worst after: {max(post_ratio):.1e}",
+ xy=(steps[-1], post_ratio[-1]), xytext=(180, 2.5e-7),
+ arrowprops={"arrowstyle": "-", "color": blue, "linewidth": 1},
+ color=blue, fontsize=8.5)
+
+ colors = [gray, orange, red, "#9A9A9A", blue]
+ bars = right.bar(range(len(accuracies)), accuracies, color=colors,
+ width=0.72)
+ right.set_xticks(range(len(labels)), labels)
+ right.tick_params(axis="x", labelsize=8.5)
+ right.set_ylim(0, 100)
+ right.set_ylabel("Validation accuracy (%)")
+ right.set_title("c The stable signal restores learning", loc="left")
+ right.grid(axis="y", color="#E1E1E1", linewidth=0.7)
+ for bar, value in zip(bars, accuracies):
+ right.text(bar.get_x() + bar.get_width() / 2, value + 2.0,
+ f"{value:.2f}", ha="center", va="bottom",
+ fontsize=8.5, fontweight="bold")
+ right.text(
+ 0.98, 0.53,
+ "20 epochs · seed 0\n1.326× BP MACs\n0 task-loss queries",
+ transform=right.transAxes, ha="right", va="center", fontsize=8.5,
+ bbox={"boxstyle": "round,pad=0.4", "facecolor": "#E8F2F8",
+ "edgecolor": "none"})
+
+ fig.suptitle(
+ "Dynamic somato-dendritic innovation turns a multiplicative failure into a stable local update",
+ y=1.02, fontsize=13.2, fontweight="bold")
+ fig.tight_layout(w_pad=2.4)
+ os.makedirs(os.path.dirname(os.path.abspath(args.out)), exist_ok=True)
+ png = args.out + ".png"
+ pdf = args.out + ".pdf"
+ fig.savefig(png, dpi=240, bbox_inches="tight", facecolor="white")
+ fig.savefig(pdf, bbox_inches="tight", metadata=PDF_METADATA,
+ facecolor="white")
+ plt.close(fig)
+
+ sources = [args.dynamic_train, args.dynamic_train_gate, args.fixed_lesion,
+ args.dynamic_short, args.dynamic_short_gate, args.mt1_gate,
+ args.kp_gate]
+ output = {
+ "script": os.path.relpath(__file__),
+ "script_sha256": sha256(__file__),
+ "sources": {path: sha256(path) for path in sources},
+ "png": png,
+ "png_sha256": sha256(png),
+ "pdf": pdf,
+ "pdf_sha256": sha256(pdf),
+ "fixed_predictor_curve_is_pre_grid_diagnostic": True,
+ }
+ manifest = args.out + "_manifest.json"
+ output["manifest"] = manifest
+ with open(manifest, "w") as handle:
+ json.dump(output, handle, indent=2, sort_keys=True)
+ handle.write("\n")
+ print(json.dumps(output, indent=2, sort_keys=True))
+
+
+if __name__ == "__main__":
+ main()