diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-29 18:43:05 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-29 18:43:05 -0500 |
| commit | e9f1342fc8e233a4841b7eb3c1363324e90ecda9 (patch) | |
| tree | 037eff7999f7abf4cd06c4f481c008529a0e45fe /experiments/plot_digital_additive_transfer.py | |
| parent | 28a8e0b249ed1847b38e116e34a1fa0efb467705 (diff) | |
figure: show SDIL transfer across digital learners
Diffstat (limited to 'experiments/plot_digital_additive_transfer.py')
| -rw-r--r-- | experiments/plot_digital_additive_transfer.py | 325 |
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() |
