summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-29 18:43:05 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-29 18:43:05 -0500
commite9f1342fc8e233a4841b7eb3c1363324e90ecda9 (patch)
tree037eff7999f7abf4cd06c4f481c008529a0e45fe /experiments
parent28a8e0b249ed1847b38e116e34a1fa0efb467705 (diff)
figure: show SDIL transfer across digital learners
Diffstat (limited to 'experiments')
-rw-r--r--experiments/plot_digital_additive_transfer.py325
1 files changed, 325 insertions, 0 deletions
diff --git a/experiments/plot_digital_additive_transfer.py b/experiments/plot_digital_additive_transfer.py
new file mode 100644
index 0000000..5e12a11
--- /dev/null
+++ b/experiments/plot_digital_additive_transfer.py
@@ -0,0 +1,325 @@
+#!/usr/bin/env python3
+"""Plot SDIL as an additive correction across four digital learners."""
+
+from __future__ import annotations
+
+import argparse
+import csv
+import json
+from pathlib import Path
+
+import matplotlib as mpl
+import matplotlib.pyplot as plt
+import numpy as np
+
+
+COLORS = {
+ "clean": "#222222",
+ "noise": "#CC79A7",
+ "raw": "#D55E00",
+ "constant": "#777777",
+ "sdil": "#0072B2",
+ "oracle": "#009E73",
+}
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument(
+ "--dualprop", type=Path,
+ default=Path("results/contrastive_bias/c1_gate.json"))
+ parser.add_argument(
+ "--ep", type=Path,
+ default=Path("results/ep_bias/c1_gate.json"))
+ parser.add_argument(
+ "--clln", type=Path,
+ default=Path("results/coupled_ladder/p2_confirm_side4.json"))
+ parser.add_argument(
+ "--overclamp", type=Path,
+ default=Path("results/coupled_ladder/p3_overclamp_side4.json"))
+ parser.add_argument(
+ "--output-analysis", type=Path,
+ default=Path("results/digital_additive_transfer.json"))
+ parser.add_argument(
+ "--output-csv", type=Path,
+ default=Path("results/digital_additive_transfer.csv"))
+ parser.add_argument(
+ "--output-figure", type=Path,
+ default=Path("results/figs/figure_digital_additive_transfer"))
+ parser.add_argument("--bootstrap-replicates", type=int, default=20000)
+ parser.add_argument("--bootstrap-seed", type=int, default=20260829)
+ return parser.parse_args()
+
+
+def summarize_values(
+ values: np.ndarray,
+ rng: np.random.Generator,
+ replicates: int,
+) -> dict:
+ values = np.asarray(values, dtype=float)
+ samples = rng.integers(0, len(values), size=(replicates, len(values)))
+ means = np.mean(values[samples], axis=1)
+ return {
+ "mean_accuracy_percent": float(np.mean(values)),
+ "bootstrap_95ci_percent": [
+ float(value) for value in np.percentile(means, (2.5, 97.5))
+ ],
+ "independent_units": len(values),
+ }
+
+
+def task_cluster_accuracies(report: dict, method: str) -> np.ndarray:
+ tasks = sorted({record["task_index"] for record in report["records"]})
+ return np.asarray([
+ 100.0 * np.mean([
+ 1.0 - record["methods"][method]["classification_error"]
+ for record in report["records"]
+ if record["task_index"] == task
+ ])
+ for task in tasks
+ ])
+
+
+def gap_recovered(methods: dict) -> float:
+ clean = methods["clean"]["mean_accuracy_percent"]
+ raw = methods["raw"]["mean_accuracy_percent"]
+ sdil = methods["sdil"]["mean_accuracy_percent"]
+ return float(100.0 * (sdil - raw) / (clean - raw))
+
+
+def build_analysis(
+ dualprop: dict,
+ ep: dict,
+ clln: dict,
+ overclamp: dict,
+ *,
+ replicates: int,
+ seed: int,
+) -> dict:
+ rng = np.random.default_rng(seed)
+ dp_rows = dualprop["rows"]
+ ep_rows = ep["rows"]
+ panels = {
+ "dualprop": {
+ "title": "Dual Propagation · CIFAR-10",
+ "unit": "seed",
+ "methods": {
+ "clean": summarize_values(np.asarray([
+ row["same_path_clean"] for row in dp_rows]), rng, replicates),
+ "raw": summarize_values(np.asarray([
+ row["raw"] for row in dp_rows]), rng, replicates),
+ "sdil": summarize_values(np.asarray([
+ row["innovation"] for row in dp_rows]), rng, replicates),
+ "oracle": summarize_values(np.asarray([
+ row["oracle"] for row in dp_rows]), rng, replicates),
+ },
+ },
+ "ep": {
+ "title": "Equilibrium Propagation · FashionMNIST",
+ "unit": "seed",
+ "methods": {
+ "clean": summarize_values(100.0 * np.asarray([
+ row["clean"] for row in ep_rows]), rng, replicates),
+ "noise": summarize_values(100.0 * np.asarray([
+ row["same_rms_noise"] for row in ep_rows]), rng, replicates),
+ "raw": summarize_values(100.0 * np.asarray([
+ row["raw"] for row in ep_rows]), rng, replicates),
+ "constant": summarize_values(100.0 * np.asarray([
+ row["constant"] for row in ep_rows]), rng, replicates),
+ "sdil": summarize_values(100.0 * np.asarray([
+ row["innovation"] for row in ep_rows]), rng, replicates),
+ "oracle": summarize_values(100.0 * np.asarray([
+ row["oracle"] for row in ep_rows]), rng, replicates),
+ },
+ },
+ "clln": {
+ "title": "Coupled learning · 32 edges",
+ "unit": "task; device draws averaged within task",
+ "methods": {
+ method: summarize_values(
+ task_cluster_accuracies(clln, source), rng, replicates)
+ for method, source in (
+ ("clean", "clean"),
+ ("noise", "matched_noise"),
+ ("raw", "raw"),
+ ("constant", "constant"),
+ ("sdil", "sdil"),
+ )
+ },
+ },
+ "overclamp": {
+ "title": "Overclamped coupled learning · 32 edges",
+ "unit": "task; device draws averaged within task",
+ "methods": {
+ method: summarize_values(
+ task_cluster_accuracies(overclamp, source), rng, replicates)
+ for method, source in (
+ ("clean", "overclamp_clean"),
+ ("raw", "overclamp"),
+ ("sdil", "overclamp_sdil"),
+ )
+ },
+ },
+ }
+ for panel in panels.values():
+ panel["raw_to_clean_gap_recovered_percent"] = gap_recovered(
+ panel["methods"])
+ return {
+ "analysis": "digital_additive_sdil_transfer",
+ "confirmatory_sources": {
+ "dualprop": dualprop["gate"] == "pass",
+ "ep": ep["gate"] == "pass",
+ "clln": bool(clln["confirmatory"]),
+ "overclamp": bool(overclamp["confirmatory"]),
+ },
+ "bootstrap": {
+ "replicates": replicates,
+ "seed": seed,
+ "interval": "percentile 95%",
+ },
+ "panels": panels,
+ }
+
+
+def write_csv(path: Path, analysis: dict) -> None:
+ path.parent.mkdir(parents=True, exist_ok=True)
+ with path.open("w", newline="") as stream:
+ writer = csv.writer(stream)
+ writer.writerow((
+ "panel", "method", "mean_accuracy_percent", "ci_low", "ci_high",
+ "independent_unit", "independent_units", "gap_recovered_percent",
+ ))
+ for panel_name, panel in analysis["panels"].items():
+ for method, summary in panel["methods"].items():
+ writer.writerow((
+ panel_name,
+ method,
+ summary["mean_accuracy_percent"],
+ *summary["bootstrap_95ci_percent"],
+ panel["unit"],
+ summary["independent_units"],
+ panel["raw_to_clean_gap_recovered_percent"],
+ ))
+
+
+def plot(path: Path, analysis: dict) -> None:
+ mpl.rcParams.update({
+ "font.family": "DejaVu Sans",
+ "font.size": 8.5,
+ "axes.labelsize": 9,
+ "axes.titlesize": 9.5,
+ "xtick.labelsize": 7.8,
+ "ytick.labelsize": 8,
+ "axes.spines.top": False,
+ "axes.spines.right": False,
+ "svg.fonttype": "none",
+ "pdf.fonttype": 42,
+ "figure.facecolor": "white",
+ "axes.facecolor": "white",
+ })
+ order = ("dualprop", "ep", "clln", "overclamp")
+ figure, axes = plt.subplots(2, 2, figsize=(8.0, 5.8), sharey=True)
+ for index, (axis, panel_name) in enumerate(zip(axes.ravel(), order)):
+ panel = analysis["panels"][panel_name]
+ methods = list(panel["methods"])
+ summaries = [panel["methods"][method] for method in methods]
+ means = np.asarray([
+ summary["mean_accuracy_percent"] for summary in summaries])
+ intervals = np.asarray([
+ summary["bootstrap_95ci_percent"] for summary in summaries])
+ positions = np.arange(len(methods))
+ axis.bar(
+ positions,
+ means,
+ color=[COLORS[method] for method in methods],
+ width=0.72,
+ edgecolor="white",
+ linewidth=0.4,
+ )
+ axis.errorbar(
+ positions,
+ means,
+ yerr=np.vstack((
+ means - intervals[:, 0], intervals[:, 1] - means,
+ )),
+ fmt="none",
+ ecolor="#222222",
+ elinewidth=0.8,
+ capsize=2.2,
+ )
+ labels = {
+ "clean": "Clean", "noise": "Same-RMS\nnoise", "raw": "Raw",
+ "constant": "Static\ncalibration", "sdil": "SDIL",
+ "oracle": "Oracle",
+ }
+ axis.set_xticks(positions, [labels[method] for method in methods])
+ axis.set_ylim(0.0, 106.0)
+ axis.grid(axis="y", color="#D9D9D9", linewidth=0.55, alpha=0.8)
+ axis.tick_params(length=3)
+ letter = chr(ord("a") + index)
+ recovery = panel["raw_to_clean_gap_recovered_percent"]
+ axis.set_title(
+ f"({letter}) {panel['title']}\n{recovery:.0f}% of clean gap recovered")
+ if index % 2 == 0:
+ axis.set_ylabel("Task accuracy (%)")
+ figure.text(
+ 0.995,
+ 0.005,
+ "Dual Prop/EP: 5 seeds; coupled learners: 40 tasks × 3 device draws",
+ ha="right",
+ va="bottom",
+ fontsize=7,
+ color="#666666",
+ )
+ figure.tight_layout(rect=(0.0, 0.035, 1.0, 1.0), h_pad=2.1, w_pad=1.8)
+ path.parent.mkdir(parents=True, exist_ok=True)
+ figure.savefig(path.with_suffix(".svg"), bbox_inches="tight")
+ figure.savefig(path.with_suffix(".pdf"), bbox_inches="tight")
+ figure.savefig(path.with_suffix(".png"), dpi=240, bbox_inches="tight")
+ plt.close(figure)
+
+
+def main() -> None:
+ args = parse_args()
+ reports = {
+ name: json.loads(path.read_text())
+ for name, path in (
+ ("dualprop", args.dualprop),
+ ("ep", args.ep),
+ ("clln", args.clln),
+ ("overclamp", args.overclamp),
+ )
+ }
+ analysis = build_analysis(
+ **reports,
+ replicates=args.bootstrap_replicates,
+ seed=args.bootstrap_seed,
+ )
+ analysis["sources"] = {
+ "dualprop": str(args.dualprop),
+ "ep": str(args.ep),
+ "clln": str(args.clln),
+ "overclamp": str(args.overclamp),
+ }
+ args.output_analysis.parent.mkdir(parents=True, exist_ok=True)
+ args.output_analysis.write_text(json.dumps(analysis, indent=2) + "\n")
+ write_csv(args.output_csv, analysis)
+ plot(args.output_figure, analysis)
+ print(json.dumps({
+ name: {
+ "gap_recovered_percent": panel[
+ "raw_to_clean_gap_recovered_percent"],
+ "methods": {
+ method: summary["mean_accuracy_percent"]
+ for method, summary in panel["methods"].items()
+ },
+ }
+ for name, panel in analysis["panels"].items()
+ }, indent=2))
+ print(f"wrote {args.output_analysis}")
+ print(f"wrote {args.output_csv}")
+ print(f"wrote {args.output_figure}.svg/.pdf/.png")
+
+
+if __name__ == "__main__":
+ main()