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