diff options
Diffstat (limited to 'experiments')
| -rwxr-xr-x | experiments/plot_dynamic_stability.py | 220 |
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() |
