summaryrefslogtreecommitdiff
path: root/experiments/physical_grid_autozero_p7.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-29 16:19:42 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-29 16:19:42 -0500
commitd9ac6e702b7671cd2c2157a1bc67c0a2a80e4638 (patch)
tree627268a452927f01899366a74cf8cc7b208ce27d /experiments/physical_grid_autozero_p7.py
parent7a81f59cc943a9893c83d5ef6b47cc50baf69584 (diff)
exp: add hardware autozero robustness matrix
Diffstat (limited to 'experiments/physical_grid_autozero_p7.py')
-rw-r--r--experiments/physical_grid_autozero_p7.py369
1 files changed, 369 insertions, 0 deletions
diff --git a/experiments/physical_grid_autozero_p7.py b/experiments/physical_grid_autozero_p7.py
new file mode 100644
index 0000000..808d0af
--- /dev/null
+++ b/experiments/physical_grid_autozero_p7.py
@@ -0,0 +1,369 @@
+#!/usr/bin/env python3
+"""Test a hardware-native auto-zero SDIL primitive on the physical grid.
+
+The local sampler observes the learning-circuit output with identical free and
+clamped edge voltages. It never reads the simulated component gains, offsets,
+or their generating distribution. The held output is subtracted during the
+subsequent learning update.
+"""
+
+from __future__ import annotations
+
+import argparse
+from concurrent.futures import ProcessPoolExecutor, as_completed
+import json
+from pathlib import Path
+import sys
+
+import numpy as np
+
+ROOT = Path(__file__).resolve().parents[1]
+sys.path.insert(0, str(ROOT))
+sys.path.insert(0, str(Path(__file__).resolve().parent))
+
+from physical_grid_bias_p5 import select_tasks # noqa: E402
+from sdil.physical_grid import ( # noqa: E402
+ AutozeroSampleHold,
+ GridCircuit,
+ GridSquareLawImperfection,
+ RingClassificationDataset,
+ train_grid_classifier,
+)
+
+
+def conditions() -> list[dict]:
+ settings = [{
+ "name": "ideal_sample_hold",
+ "sample_gain": 1.0,
+ "pedestal_standard_deviation_v_per_s": 0.0,
+ "sample_noise_standard_deviation_v_per_s": 0.0,
+ "refresh_interval_updates": 1,
+ "method": "autozero_sdil",
+ }]
+ for noise in (0.1, 0.25, 0.5, 1.0):
+ settings.append({
+ "name": f"sample_noise_{noise:g}",
+ "sample_gain": 1.0,
+ "pedestal_standard_deviation_v_per_s": 0.0,
+ "sample_noise_standard_deviation_v_per_s": noise,
+ "refresh_interval_updates": 1,
+ "method": "autozero_sdil",
+ })
+ for pedestal in (0.1, 0.25, 0.5, 1.0):
+ settings.append({
+ "name": f"pedestal_{pedestal:g}",
+ "sample_gain": 1.0,
+ "pedestal_standard_deviation_v_per_s": pedestal,
+ "sample_noise_standard_deviation_v_per_s": 0.0,
+ "refresh_interval_updates": 1,
+ "method": "autozero_sdil",
+ })
+ for gain in (0.9, 0.95, 0.99, 1.01, 1.05, 1.1):
+ settings.append({
+ "name": f"sample_gain_{gain:g}",
+ "sample_gain": gain,
+ "pedestal_standard_deviation_v_per_s": 0.0,
+ "sample_noise_standard_deviation_v_per_s": 0.0,
+ "refresh_interval_updates": 1,
+ "method": "autozero_sdil",
+ })
+ for interval in (2, 4, 8, 16):
+ settings.append({
+ "name": f"refresh_every_{interval}",
+ "sample_gain": 1.0,
+ "pedestal_standard_deviation_v_per_s": 0.0,
+ "sample_noise_standard_deviation_v_per_s": 0.0,
+ "refresh_interval_updates": interval,
+ "method": "autozero_sdil",
+ })
+ settings.extend((
+ {
+ "name": "combined_mild",
+ "sample_gain": 0.99,
+ "pedestal_standard_deviation_v_per_s": 0.1,
+ "sample_noise_standard_deviation_v_per_s": 0.1,
+ "refresh_interval_updates": 1,
+ "method": "autozero_sdil",
+ },
+ {
+ "name": "combined_strong",
+ "sample_gain": 0.95,
+ "pedestal_standard_deviation_v_per_s": 0.25,
+ "sample_noise_standard_deviation_v_per_s": 0.25,
+ "refresh_interval_updates": 1,
+ "method": "autozero_sdil",
+ },
+ {
+ "name": "overclamp_plus_ideal_sample_hold",
+ "sample_gain": 1.0,
+ "pedestal_standard_deviation_v_per_s": 0.0,
+ "sample_noise_standard_deviation_v_per_s": 0.0,
+ "refresh_interval_updates": 1,
+ "method": "overclamp_autozero_sdil",
+ },
+ {
+ "name": "overclamp_plus_combined_mild",
+ "sample_gain": 0.99,
+ "pedestal_standard_deviation_v_per_s": 0.1,
+ "sample_noise_standard_deviation_v_per_s": 0.1,
+ "refresh_interval_updates": 1,
+ "method": "overclamp_autozero_sdil",
+ },
+ ))
+ return settings
+
+
+def run_job(job: dict) -> dict:
+ circuit = GridCircuit()
+ task = job["task"]
+ condition = job["condition"]
+ task_index = job["task_index"]
+ device_seed = job["device_seed"]
+ dataset = RingClassificationDataset(
+ inputs_v=np.asarray(task["inputs_v"], dtype=float).T,
+ labels_v=(
+ 2.0 * np.asarray(task["classes"], dtype=float) - 1.0
+ ) * 0.018,
+ )
+ initial_gates = np.asarray(task["initial_gates_v"], dtype=float)
+ imperfection = GridSquareLawImperfection.sample_appendix_c(
+ circuit.edge_count, device_seed)
+ sampler_seed = device_seed + 5_000_003 + task_index
+ pedestal_rng = np.random.default_rng(
+ device_seed + 4_000_003 + task_index)
+ pedestal = pedestal_rng.normal(
+ 0.0,
+ condition["pedestal_standard_deviation_v_per_s"],
+ circuit.edge_count,
+ )
+ sampler = AutozeroSampleHold(
+ sample_gain=condition["sample_gain"],
+ pedestal_offset_v_per_s=pedestal,
+ sample_noise_standard_deviation_v_per_s=(
+ condition["sample_noise_standard_deviation_v_per_s"]),
+ refresh_interval_updates=condition["refresh_interval_updates"],
+ )
+ kwargs = dict(
+ circuit=circuit,
+ initial_gates=initial_gates,
+ dataset=dataset,
+ imperfection=imperfection,
+ method=condition["method"],
+ autozero_sample_hold=sampler,
+ autozero_seed=sampler_seed,
+ )
+ try:
+ if condition["method"].startswith("overclamp"):
+ result = train_grid_classifier(
+ **kwargs,
+ epochs=job["overclamp_epochs"],
+ overclamp_time_seconds_per_v=0.0025,
+ record_every=10,
+ early_stop_perfect_checkpoints=3,
+ )
+ else:
+ result = train_grid_classifier(
+ **kwargs,
+ epochs=job["standard_epochs"],
+ standard_learning_time_seconds=1e-3,
+ record_every=50,
+ )
+ except RuntimeError as error:
+ return {
+ "condition": condition["name"],
+ "task_index": task_index,
+ "device_seed": device_seed,
+ "input_diameter_v": task["input_diameter_v"],
+ "status": "circuit_solver_failure",
+ "failure_message": str(error),
+ "classification_error": 1.0,
+ "zero_error": False,
+ "autozero_samples": None,
+ "local_updates": None,
+ "autozero_sample_fraction_per_update": None,
+ "autozero_baseline_rmse_v_per_s": None,
+ "max_abs_clamp_displacement_v": None,
+ }
+ return {
+ "condition": condition["name"],
+ "task_index": task_index,
+ "device_seed": device_seed,
+ "input_diameter_v": task["input_diameter_v"],
+ "status": "completed",
+ "classification_error": result["classification_error"],
+ "zero_error": result["classification_error"] == 0.0,
+ "hinge_loss_v2": result["hinge_loss_v2"],
+ "autozero_samples": result["autozero_samples"],
+ "local_updates": result["local_updates"],
+ "autozero_sample_fraction_per_update": (
+ result["autozero_sample_fraction_per_update"]),
+ "autozero_baseline_rmse_v_per_s": (
+ result["autozero_baseline_rmse_v_per_s"]),
+ "max_abs_clamp_displacement_v": (
+ result["max_abs_clamp_displacement_v"]),
+ }
+
+
+def reference_records(path: Path) -> dict[tuple[int, int], dict]:
+ report = json.loads(path.read_text())
+ return {
+ (record["task_index"], record["device_seed"]): record
+ for record in report["records"]
+ }
+
+
+def summarize(
+ records: list[dict], settings: list[dict], reference: dict
+) -> dict:
+ summaries = {}
+ for setting in settings:
+ selected = [
+ record for record in records
+ if record["condition"] == setting["name"]
+ ]
+ errors = np.asarray([
+ record["classification_error"] for record in selected
+ ])
+ completed = [
+ record for record in selected if record["status"] == "completed"
+ ]
+ overclamp_errors = np.asarray([
+ reference[(record["task_index"], record["device_seed"])]
+ ["methods"]["overclamp"]["classification_error"]
+ for record in selected
+ ])
+ summaries[setting["name"]] = {
+ **setting,
+ "trials": len(selected),
+ "mean_classification_error": float(np.mean(errors)),
+ "median_classification_error": float(np.median(errors)),
+ "zero_error_fraction": float(np.mean(errors == 0.0)),
+ "solver_failure_fraction": float(np.mean([
+ record["status"] != "completed" for record in selected
+ ])),
+ "lower_error_than_overclamp_fraction": float(np.mean(
+ errors < overclamp_errors)),
+ "equal_error_to_overclamp_fraction": float(np.mean(
+ errors == overclamp_errors)),
+ "higher_error_than_overclamp_fraction": float(np.mean(
+ errors > overclamp_errors)),
+ "median_autozero_samples": float(np.median([
+ record["autozero_samples"] for record in completed
+ ])),
+ "mean_autozero_sample_fraction_per_update": float(np.mean([
+ record["autozero_sample_fraction_per_update"]
+ for record in completed
+ ])),
+ "median_autozero_baseline_rmse_v_per_s": float(np.median([
+ record["autozero_baseline_rmse_v_per_s"]
+ for record in completed
+ ])),
+ "mean_max_abs_clamp_displacement_v": float(np.mean([
+ record["max_abs_clamp_displacement_v"]
+ for record in completed
+ ])),
+ }
+ return summaries
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument(
+ "--protocol", type=Path,
+ default=Path("results/physical_bias/dillavou_fig5_protocol.json"))
+ parser.add_argument(
+ "--reference", type=Path,
+ default=Path(
+ "results/physical_bias/p5_full_grid_bias_crossover.json"))
+ parser.add_argument(
+ "--output", type=Path,
+ default=Path(
+ "results/physical_bias/p7_grid_autozero_robustness.json"))
+ parser.add_argument("--rotations", type=int, default=8)
+ parser.add_argument(
+ "--device-seeds", default="20260829,20260830,20260831,20260832")
+ parser.add_argument("--standard-epochs", type=int, default=600)
+ parser.add_argument("--overclamp-epochs", type=int, default=1000)
+ parser.add_argument("--workers", type=int, default=16)
+ parser.add_argument("--conditions", default="all")
+ return parser.parse_args()
+
+
+def main() -> None:
+ args = parse_args()
+ protocol = json.loads(args.protocol.read_text())
+ tasks = select_tasks(protocol, args.rotations)
+ device_seeds = tuple(int(seed) for seed in args.device_seeds.split(","))
+ all_settings = conditions()
+ if args.conditions == "all":
+ settings = all_settings
+ else:
+ requested = set(args.conditions.split(","))
+ settings = [
+ setting for setting in all_settings
+ if setting["name"] in requested
+ ]
+ missing = requested - {setting["name"] for setting in settings}
+ if missing:
+ raise ValueError(f"unknown conditions: {sorted(missing)}")
+ jobs = []
+ for condition in settings:
+ for task_index, task in enumerate(tasks):
+ for device_seed in device_seeds:
+ jobs.append({
+ "condition": condition,
+ "task_index": task_index,
+ "task": task,
+ "device_seed": device_seed,
+ "standard_epochs": args.standard_epochs,
+ "overclamp_epochs": args.overclamp_epochs,
+ })
+ records = []
+ with ProcessPoolExecutor(max_workers=args.workers) as executor:
+ futures = [executor.submit(run_job, job) for job in jobs]
+ for completed, future in enumerate(as_completed(futures), start=1):
+ records.append(future.result())
+ if completed % 80 == 0 or completed == len(jobs):
+ print(f"completed {completed}/{len(jobs)}", flush=True)
+ records.sort(key=lambda record: (
+ record["condition"], record["task_index"], record["device_seed"]))
+ references = reference_records(args.reference)
+ report = {
+ "analysis": "physical_grid_hardware_autozero_p7",
+ "confirmatory": False,
+ "autodiff_used": False,
+ "source_protocol": str(args.protocol),
+ "reference_results": str(args.reference),
+ "protocol": {
+ "task_count": len(tasks),
+ "rotations_per_input_diameter": args.rotations,
+ "device_seeds": device_seeds,
+ "trials_per_condition": len(tasks) * len(device_seeds),
+ "component_imperfection": {
+ "measurement_gain_standard_deviation": 0.01,
+ "twin_input_mismatch_standard_deviation_v": 0.001,
+ "multiplier_output_offset_standard_deviation_v_per_s": 2.3,
+ },
+ "autozero_operation": (
+ "Each edge samples its own raw learning-circuit output with "
+ "free=clamped, holds that voltage, then subtracts it from the "
+ "active update. The rule does not read component parameters."),
+ "sample_count_scope": (
+ "One parallel network-wide sample event contains one local "
+ "sample on every edge."),
+ "standard_epochs": args.standard_epochs,
+ "standard_learning_time_seconds": 1e-3,
+ "overclamp_epochs_maximum": args.overclamp_epochs,
+ "conditions": settings,
+ },
+ "records": records,
+ "summary": summarize(records, settings, references),
+ }
+ args.output.parent.mkdir(parents=True, exist_ok=True)
+ args.output.write_text(json.dumps(report, indent=2) + "\n")
+ print(json.dumps(report["summary"], indent=2))
+ print(f"wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()