#!/usr/bin/env python3 """Evaluate correlated auto-zero sampling on the physical CLLN grid.""" 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 CorrelatedDoubleSampleHold, GridCircuit, GridSquareLawImperfection, RingClassificationDataset, train_grid_classifier, ) def setting( name: str, *, common_pedestal: float = 0.0, pedestal_mismatch: float = 0.0, gain_mismatch: float = 0.0, noise: float = 0.0, refresh: int = 1, method: str = "cds_autozero_sdil", ) -> dict: return { "name": name, "common_pedestal_standard_deviation_v_per_s": common_pedestal, "pedestal_mismatch_standard_deviation_v_per_s": pedestal_mismatch, "sample_gain_mismatch_standard_deviation": gain_mismatch, "sample_noise_standard_deviation_v_per_s": noise, "refresh_interval_updates": refresh, "method": method, } def conditions() -> list[dict]: settings = [setting("ideal_cds")] for common_pedestal in (2.3, 10.0): settings.append(setting( f"common_pedestal_{common_pedestal:g}", common_pedestal=common_pedestal, )) for mismatch in (0.01, 0.025, 0.05, 0.1, 0.25, 0.5): settings.append(setting( f"pedestal_mismatch_{mismatch:g}", common_pedestal=2.3, pedestal_mismatch=mismatch, )) for mismatch in (0.001, 0.005, 0.01, 0.025, 0.05): settings.append(setting( f"gain_mismatch_{mismatch:g}", common_pedestal=2.3, gain_mismatch=mismatch, )) for noise in (0.1, 0.25, 0.5, 1.0): settings.append(setting( f"sample_noise_{noise:g}", common_pedestal=2.3, noise=noise, )) for refresh in (2, 4, 8): settings.append(setting( f"refresh_every_{refresh}", common_pedestal=2.3, refresh=refresh, )) settings.extend(( setting( "combined_mild", common_pedestal=2.3, pedestal_mismatch=0.025, gain_mismatch=0.005, noise=0.1, ), setting( "combined_mild_refresh4", common_pedestal=2.3, pedestal_mismatch=0.025, gain_mismatch=0.005, noise=0.1, refresh=4, ), setting( "combined_strong", common_pedestal=2.3, pedestal_mismatch=0.1, gain_mismatch=0.01, noise=0.25, ), setting( "overclamp_plus_combined_mild", common_pedestal=2.3, pedestal_mismatch=0.025, gain_mismatch=0.005, noise=0.1, method="overclamp_cds_autozero_sdil", ), )) return settings def fixed_edge_errors( circuit: GridCircuit, task_index: int, device_seed: int, condition: dict ) -> tuple[np.ndarray, np.ndarray, np.ndarray]: common_rng = np.random.default_rng( device_seed + 6_000_003 + task_index) pedestal_rng = np.random.default_rng( device_seed + 7_000_003 + task_index) gain_rng = np.random.default_rng( device_seed + 8_000_003 + task_index) return ( common_rng.normal( 0.0, condition["common_pedestal_standard_deviation_v_per_s"], circuit.edge_count, ), pedestal_rng.normal( 0.0, condition["pedestal_mismatch_standard_deviation_v_per_s"], circuit.edge_count, ), gain_rng.normal( 0.0, condition["sample_gain_mismatch_standard_deviation"], circuit.edge_count, ), ) 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, ) imperfection = GridSquareLawImperfection.sample_appendix_c( circuit.edge_count, device_seed) common_pedestal, pedestal_mismatch, gain_mismatch = fixed_edge_errors( circuit, task_index, device_seed, condition) sampler = CorrelatedDoubleSampleHold( common_sample_gain=1.0, sample_gain_mismatch=gain_mismatch, common_pedestal_offset_v_per_s=common_pedestal, pedestal_mismatch_v_per_s=pedestal_mismatch, 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=np.asarray(task["initial_gates_v"], dtype=float), dataset=dataset, imperfection=imperfection, method=condition["method"], correlated_sample_hold=sampler, autozero_seed=device_seed + 9_000_003 + task_index, ) 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, "neutral_samples": None, "active_samples": None, "local_updates": None, "applied_rate_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"], "neutral_samples": result["autozero_samples"], "active_samples": result["autozero_active_samples"], "local_updates": result["local_updates"], "neutral_sample_fraction_per_update": ( result["autozero_sample_fraction_per_update"]), "applied_rate_rmse_v_per_s": ( result["autozero_applied_rate_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: output = {} for condition in settings: selected = [ record for record in records if record["condition"] == condition["name"] ] completed = [ record for record in selected if record["status"] == "completed" ] errors = np.asarray([ record["classification_error"] for record in selected ]) overclamp_errors = np.asarray([ reference[(record["task_index"], record["device_seed"])] ["methods"]["overclamp"]["classification_error"] for record in selected ]) output[condition["name"]] = { **condition, "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)), "mean_neutral_sample_fraction_per_update": float(np.mean([ record["neutral_sample_fraction_per_update"] for record in completed ])), "median_applied_rate_rmse_v_per_s": float(np.median([ record["applied_rate_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 output 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/p9_grid_correlated_autozero.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 = [ condition for condition in all_settings if condition["name"] in requested ] missing = requested - {condition["name"] for condition in settings} if missing: raise ValueError(f"unknown conditions: {sorted(missing)}") jobs = [{ "condition": condition, "task_index": task_index, "task": task, "device_seed": device_seed, "standard_epochs": args.standard_epochs, "overclamp_epochs": args.overclamp_epochs, } for condition in settings for task_index, task in enumerate(tasks) for device_seed in device_seeds] 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_correlated_autozero_p9", "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, }, "sampling_operation": ( "Each edge samples neutral and active outputs through matched " "local paths and applies their difference. Common sample-path " "pedestal cancels without reading component parameters."), "standard_epochs": args.standard_epochs, "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()