#!/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()